FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

GitHub Viewer

#include #include "model.hpp" #include "constants.hpp" #include "tile_bounds.hpp" #include "project_gaussians.hpp" #include "rasterize_gaussians.hpp" #include "tensor_math.hpp" #include "gsplat.hpp" #include "utils.hpp" #ifdef USE_MPS #include #endif #ifdef USE_HIP #include #elif defined(USE_CUDA) #include #endif namespace fs = std::filesystem; torch::Tensor randomQuatTensor(long long n){ torch::Tensor u = torch::rand(n); torch::Tensor v = torch::rand(n); torch::Tensor w = torch::rand(n); return torch::stack({ torch::sqrt(1 - u) * torch::sin(2 * PI * v), torch::sqrt(1 - u) * torch::cos(2 * PI * v), torch::sqrt(u) * torch::sin(2 * PI * w), torch::sqrt(u) * torch::cos(2 * PI * w) }, -1); } torch::Tensor projectionMatrix(float zNear, float zFar, float fovX, float fovY, const torch::Device &device){ // OpenGL perspective projection matrix float t = zNear * std::tan(0.5f * fovY); float b = -t; float r = zNear * std::tan(0.5f * fovX); float l = -r; return torch::tensor({ {2.0f * zNear / (r - l), 0.0f, (r + l) / (r - l), 0.0f}, {0.0f, 2 * zNear / (t - b), (t + b) / (t - b), 0.0f}, {0.0f, 0.0f, (zFar + zNear) / (zFar - zNear), -1.0f * zFar * zNear / (zFar - zNear)}, {0.0f, 0.0f, 1.0f, 0.0f} }, device); } torch::Tensor psnr(const torch::Tensor& rendered, const torch::Tensor& gt){ torch::Tensor mse = (rendered - gt).pow(2).mean(); return (10.f * torch::log10(1.0 / mse)); } torch::Tensor l1(const torch::Tensor& rendered, const torch::Tensor& gt){ return torch::abs(gt - rendered).mean(); } void Model::setupOptimizers(){ releaseOptimizers(); meansOpt = new torch::optim::Adam({means}, torch::optim::AdamOptions(0.00016)); scalesOpt = new torch::optim::Adam({scales}, torch::optim::AdamOptions(0.005)); quatsOpt = new torch::optim::Adam({quats}, torch::optim::AdamOptions(0.001)); featuresDcOpt = new torch::optim::Adam({featuresDc}, torch::optim::AdamOptions(0.0025)); featuresRestOpt = new torch::optim::Adam({featuresRest}, torch::optim::AdamOptions(0.000125)); opacitiesOpt = new torch::optim::Adam({opacities}, torch::optim::AdamOptions(0.05)); meansOptScheduler = new OptimScheduler(meansOpt, 0.0000016f, maxSteps); } void Model::releaseOptimizers(){ RELEASE_SAFELY(meansOpt); RELEASE_SAFELY(scalesOpt); RELEASE_SAFELY(quatsOpt); RELEASE_SAFELY(featuresDcOpt); RELEASE_SAFELY(featuresRestOpt); RELEASE_SAFELY(opacitiesOpt); RELEASE_SAFELY(meansOptScheduler); } torch::Tensor Model::forward(Camera& cam, int step){ const float scaleFactor = getDownscaleFactor(step); const float fx = cam.fx / scaleFactor; const float fy = cam.fy / scaleFactor; const float cx = cam.cx / scaleFactor; const float cy = cam.cy / scaleFactor; const int height = static_cast(static_cast(cam.height) / scaleFactor); const int width = static_cast(static_cast(cam.width) / scaleFactor); torch::Tensor R = cam.camToWorld.index({Slice(None, 3), Slice(None, 3)}); torch::Tensor T = cam.camToWorld.index({Slice(None, 3), Slice(3,4)}); // Flip the z and y axes to align with gsplat conventions R = torch::matmul(R, torch::diag(torch::tensor({1.0f, -1.0f, -1.0f}, R.device()))); // worldToCam torch::Tensor Rinv = R.transpose(0, 1); torch::Tensor Tinv = torch::matmul(-Rinv, T); lastHeight = height; lastWidth = width; torch::Tensor viewMat = torch::eye(4, device); viewMat.index_put_({Slice(None, 3), Slice(None, 3)}, Rinv); viewMat.index_put_({Slice(None, 3), Slice(3, 4)}, Tinv); float fovX = 2.0f * std::atan(width / (2.0f * fx)); float fovY = 2.0f * std::atan(height / (2.0f * fy)); torch::Tensor projMat = projectionMatrix(0.001f, 1000.0f, fovX, fovY, device); torch::Tensor colors = torch::cat({featuresDc.index({Slice(), None, Slice()}), featuresRest}, 1); torch::Tensor conics; torch::Tensor depths; // GPU-only torch::Tensor numTilesHit; // GPU-only torch::Tensor cov2d; // CPU-only torch::Tensor camDepths; // CPU-only torch::Tensor rgb; if (device == torch::kCPU){ auto p = ProjectGaussiansCPU::apply(means, torch::exp(scales), 1, quats / quats.norm(2, {-1}, true), viewMat, torch::matmul(projMat, viewMat), fx, fy, cx, cy, height, width); xys = p[0]; radii = p[1]; conics = p[2]; cov2d = p[3]; camDepths = p[4]; }else{ #if defined(USE_HIP) || defined(USE_CUDA) || defined(USE_MPS) TileBounds tileBounds = std::make_tuple((width + BLOCK_X - 1) / BLOCK_X, (height + BLOCK_Y - 1) / BLOCK_Y, 1); auto p = ProjectGaussians::apply(means, torch::exp(scales), 1, quats / quats.norm(2, {-1}, true), viewMat, torch::matmul(projMat, viewMat), fx, fy, cx, cy, height, width, tileBounds); xys = p[0]; depths = p[1]; radii = p[2]; conics = p[3]; numTilesHit = p[4]; #else throw std::runtime_error("GPU support not built, use --cpu"); #endif } xys.retain_grad(); if (radii.sum().item() == 0.0f) return backgroundColor.repeat({height, width, 1}); torch::Tensor viewDirs = means.detach() - T.transpose(0, 1).to(device); viewDirs = viewDirs / viewDirs.norm(2, {-1}, true); int degreesToUse = (std::min)(step / shDegreeInterval, shDegree); torch::Tensor rgbs; if (device == torch::kCPU){ rgbs = SphericalHarmonicsCPU::apply(degreesToUse, viewDirs, colors); }else{ #if defined(USE_HIP) || defined(USE_CUDA) || defined(USE_MPS) #ifdef USE_MPS torch::mps::synchronize(); #endif rgbs = SphericalHarmonics::apply(degreesToUse, viewDirs, colors); #endif } rgbs = torch::clamp_min(rgbs + 0.5f, 0.0f); if (device == torch::kCPU){ rgb = RasterizeGaussiansCPU::apply( xys, radii, conics, rgbs, torch::sigmoid(opacities), cov2d, camDepths, height, width, backgroundColor); }else{ #if defined(USE_HIP) || defined(USE_CUDA) || defined(USE_MPS) rgb = RasterizeGaussians::apply( xys, depths, radii, conics, numTilesHit, rgbs, torch::sigmoid(opacities), height, width, backgroundColor); #endif } rgb = torch::clamp_max(rgb, 1.0f); return rgb; } void Model::optimizersZeroGrad(){ meansOpt->zero_grad(); scalesOpt->zero_grad(); quatsOpt->zero_grad(); featuresDcOpt->zero_grad(); featuresRestOpt->zero_grad(); opacitiesOpt->zero_grad(); } void Model::optimizersStep(){ meansOpt->step(); scalesOpt->step(); quatsOpt->step(); featuresDcOpt->step(); featuresRestOpt->step(); opacitiesOpt->step(); } void Model::schedulersStep(int step){ meansOptScheduler->step(step); } int Model::getDownscaleFactor(int step){ return std::pow(2, (std::max)(numDownscales - step / resolutionSchedule, 0)); } void Model::addToOptimizer(torch::optim::Adam *optimizer, const torch::Tensor &newParam, const torch::Tensor &idcs, int nSamples){ torch::Tensor param = optimizer->param_groups()[0].params()[0]; #if TORCH_VERSION_MAJOR == 2 && TORCH_VERSION_MINOR > 1 auto pId = param.unsafeGetTensorImpl(); #else auto pId = c10::guts::to_string(param.unsafeGetTensorImpl()); #endif auto paramState = std::make_unique(static_cast(*optimizer->state()[pId])); std::vector repeats; repeats.push_back(nSamples); for (long int i = 0; i < paramState->exp_avg().dim() - 1; i++){ repeats.push_back(1); } paramState->exp_avg(torch::cat({ paramState->exp_avg(), torch::zeros_like(paramState->exp_avg().index({idcs.squeeze()})).repeat(repeats) }, 0)); paramState->exp_avg_sq(torch::cat({ paramState->exp_avg_sq(), torch::zeros_like(paramState->exp_avg_sq().index({idcs.squeeze()})).repeat(repeats) }, 0)); optimizer->state().erase(pId); #if TORCH_VERSION_MAJOR == 2 && TORCH_VERSION_MINOR > 1 auto newPId = newParam.unsafeGetTensorImpl(); #else auto newPId = c10::guts::to_string(newParam.unsafeGetTensorImpl()); #endif optimizer->state()[newPId] = std::move(paramState); optimizer->param_groups()[0].params()[0] = newParam; } void Model::removeFromOptimizer(torch::optim::Adam *optimizer, const torch::Tensor &newParam, const torch::Tensor &deletedMask){ torch::Tensor param = optimizer->param_groups()[0].params()[0]; #if TORCH_VERSION_MAJOR == 2 && TORCH_VERSION_MINOR > 1 auto pId = param.unsafeGetTensorImpl(); #else auto pId = c10::guts::to_string(param.unsafeGetTensorImpl()); #endif auto paramState = std::make_unique(static_cast(*optimizer->state()[pId])); paramState->exp_avg(paramState->exp_avg().index({~deletedMask})); paramState->exp_avg_sq(paramState->exp_avg_sq().index({~deletedMask})); optimizer->state().erase(pId); #if TORCH_VERSION_MAJOR == 2 && TORCH_VERSION_MINOR > 1 auto newPId = newParam.unsafeGetTensorImpl(); #else auto newPId = c10::guts::to_string(newParam.unsafeGetTensorImpl()); #endif optimizer->param_groups()[0].params()[0] = newParam; optimizer->state()[newPId] = std::move(paramState); } void Model::afterTrain(int step){ torch::NoGradGuard noGrad; // When radii.sum() == 0 if (!xys.grad().defined()) return; if (step < stopSplitAt){ torch::Tensor visibleMask = (radii > 0).flatten(); torch::Tensor grads = torch::linalg_vector_norm(xys.grad().detach(), 2, { -1 }, false, torch::kFloat32); if (!xysGradNorm.numel()){ xysGradNorm = grads; visCounts = torch::ones_like(xysGradNorm); }else{ visCounts.index_put_({visibleMask}, visCounts.index({visibleMask}) + 1); xysGradNorm.index_put_({visibleMask}, grads.index({visibleMask}) + xysGradNorm.index({visibleMask})); } if (!max2DSize.numel()){ max2DSize = torch::zeros_like(radii, torch::kFloat32); } torch::Tensor newRadii = radii.detach().index({visibleMask}); max2DSize.index_put_({visibleMask}, torch::maximum( max2DSize.index({visibleMask}), newRadii / static_cast( (std::max)(lastHeight, lastWidth) ) )); } if (step % refineEvery == 0 && step > warmupLength){ int resetInterval = resetAlphaEvery * refineEvery; bool doDensification = step < stopSplitAt && step % resetInterval > numCameras + refineEvery; torch::Tensor splitsMask; const float cullAlphaThresh = 0.1f; if (doDensification){ int numPointsBefore = means.size(0); torch::Tensor avgGradNorm = (xysGradNorm / visCounts) * 0.5f * static_cast( (std::max)(lastWidth, lastHeight) ); torch::Tensor highGrads = (avgGradNorm > densifyGradThresh).squeeze(); // Split gaussians that are too large torch::Tensor splits = (std::get(scales.exp().max(-1)) > densifySizeThresh).squeeze(); if (step < stopScreenSizeAt){ splits |= (max2DSize > splitScreenSize).squeeze(); } splits &= highGrads; const int nSplitSamples = 2; int nSplits = splits.sum().item(); torch::Tensor centeredSamples = torch::randn({nSplitSamples * nSplits, 3}, device); // Nx3 of axis-aligned scales torch::Tensor scaledSamples = torch::exp(scales.index({splits}).repeat({nSplitSamples, 1})) * centeredSamples; torch::Tensor qs = quats.index({splits}) / torch::linalg_vector_norm(quats.index({splits}), 2, { -1 }, true, torch::kFloat32); torch::Tensor rots = quatToRotMat(qs.repeat({nSplitSamples, 1})); torch::Tensor rotatedSamples = torch::bmm(rots, scaledSamples.index({"...", None})).squeeze(); torch::Tensor splitMeans = rotatedSamples + means.index({splits}).repeat({nSplitSamples, 1}); torch::Tensor splitFeaturesDc = featuresDc.index({splits}).repeat({nSplitSamples, 1}); torch::Tensor splitFeaturesRest = featuresRest.index({splits}).repeat({nSplitSamples, 1, 1}); torch::Tensor splitOpacities = opacities.index({splits}).repeat({nSplitSamples, 1}); const float sizeFac = 1.6f; torch::Tensor splitScales = torch::log(torch::exp(scales.index({splits})) / sizeFac).repeat({nSplitSamples, 1}); scales.index({splits}) = torch::log(torch::exp(scales.index({splits})) / sizeFac); torch::Tensor splitQuats = quats.index({splits}).repeat({nSplitSamples, 1}); // Duplicate gaussians that are too small torch::Tensor dups = (std::get(scales.exp().max(-1)) exp_avg_sq())); std::cout

Back | FazBrowse Home | New Git URL