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