Repository navigation
Expand file tree
/
Copy pathrenderer.cpp
More file actions
211 lines (179 loc) · 9.34 KB
/
Copy pathrenderer.cpp
File metadata and controls
211 lines (179 loc) · 9.34 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
#include "renderer.h"
#include <algorithm>
#include <tuple>
#include <vector>
using namespace torch::indexing;
namespace {
torch::Tensor head(const torch::Tensor &t) { return t.index({"...", Slice(None, 1)}); }
torch::Tensor drop_last(const torch::Tensor &t) { return t.index({"...", Slice(None, -1)}); }
torch::Tensor drop_first(const torch::Tensor &t) { return t.index({"...", Slice(1, None)}); }
// Alpha compositing along the last axis: the shared core of both render passes.
torch::Tensor weights_from_sigma(const torch::Tensor &sigma, const torch::Tensor &z_vals) {
auto dists = torch::cat(
{drop_first(z_vals) - drop_last(z_vals), torch::full_like(head(z_vals), 1e10)}, -1);
auto alpha = 1.0 - torch::exp(-sigma * dists);
auto transmittance = torch::cumprod(1.0 - alpha + 1e-10, -1);
return alpha * torch::cat({torch::ones_like(head(alpha)), drop_last(transmittance)}, -1);
}
// Midpoints between adjacent samples, used as interior bin edges.
torch::Tensor midpoints(const torch::Tensor &z) {
return 0.5 * (drop_first(z) + drop_last(z));
}
} // namespace
NeRFRenderer::NeRFRenderer(SirenNeRF &model, float focal, torch::Device device,
torch::Tensor bg_color)
: model_(model), device_(device), focal_(focal), bg_color_(bg_color.to(device)) {}
std::pair<torch::Tensor, torch::Tensor> NeRFRenderer::get_rays(
int H, int W, const torch::Tensor &pose) const {
auto opts = torch::dtype(torch::kFloat32).device(device_);
auto grid = torch::meshgrid({torch::arange(W, opts), torch::arange(H, opts)}, "xy");
// Pixel centres on the image plane; y and z are negated because the camera looks
// down -z with y up (Blender/OpenGL convention).
auto dirs = torch::stack({(grid[0] - W * 0.5f) / focal_, -(grid[1] - H * 0.5f) / focal_,
-torch::ones_like(grid[0])},
-1);
auto rays_d = (dirs.unsqueeze(-2) * pose.index({Slice(0, 3), Slice(0, 3)})).sum(-1);
auto rays_o = pose.index({Slice(0, 3), -1}).expand(rays_d.sizes());
return {rays_o.reshape({-1, 3}), rays_d.reshape({-1, 3})};
}
torch::Tensor NeRFRenderer::sample_z(int64_t n_rays, int n_samples,
const RenderOptions &opt, bool jitter) const {
auto z = torch::linspace(opt.z_near, opt.z_far, n_samples, device_)
.expand({n_rays, n_samples})
.contiguous();
if (!jitter) return z;
// NeRF 5.2: split the interval into bins and draw one sample uniformly per bin.
auto mids = midpoints(z);
auto upper = torch::cat({mids, z.index({"...", Slice(-1, None)})}, -1);
auto lower = torch::cat({head(z), mids}, -1);
return lower + (upper - lower) * torch::rand_like(z);
}
RenderOutput NeRFRenderer::render(const torch::Tensor &rays_o, const torch::Tensor &rays_d,
const RenderOptions &opt) const {
const bool jitter = opt.strategy == SampleStrategy::STRATIFIED ||
(opt.strategy == SampleStrategy::PROPOSAL && !opt.deterministic);
auto z_coarse = sample_z(rays_o.size(0), opt.n_samples, opt, jitter);
if (opt.strategy != SampleStrategy::PROPOSAL || opt.n_importance <= 0) {
auto out = volume_render(rays_o, rays_d, z_coarse, opt.batch_size);
out.fine_z = z_coarse;
return out;
}
// The coarse pass only decides where to place fine samples, so it needs no
// gradients. use_proposal=false reuses the full model instead (warm-up).
torch::Tensor z_all;
{
torch::NoGradGuard no_grad;
auto w = opt.use_proposal
? proposal_weights(rays_o, rays_d, z_coarse, opt.batch_size)
: volume_render(rays_o, rays_d, z_coarse, opt.batch_size).weights;
auto z_fine = sample_pdf(midpoints(z_coarse), w.index({"...", Slice(1, -1)}),
opt.n_importance, opt.deterministic);
z_all = std::get<0>(torch::sort(torch::cat({z_coarse, z_fine}, -1), -1));
}
auto out = volume_render(rays_o, rays_d, z_all, opt.batch_size);
out.fine_z = z_all;
out.coarse_z = z_coarse;
return out;
}
RenderOutput NeRFRenderer::render_image(int H, int W, const torch::Tensor &pose,
const RenderOptions &opt) const {
auto [rays_o, rays_d] = get_rays(H, W, pose);
auto out = render(rays_o, rays_d, opt);
out.rgb = out.rgb.view({H, W, 3});
out.depth = out.depth.view({H, W});
return out;
}
torch::Tensor NeRFRenderer::proposal_weights(const torch::Tensor &rays_o,
const torch::Tensor &rays_d,
const torch::Tensor &z_vals,
int batch_size) const {
auto pts = (rays_o.unsqueeze(-2) + rays_d.unsqueeze(-2) * z_vals.unsqueeze(-1))
.reshape({-1, 3});
std::vector<torch::Tensor> chunks;
for (int64_t i = 0; i < pts.size(0); i += batch_size) {
const auto end = std::min<int64_t>(i + batch_size, pts.size(0));
chunks.push_back(model_.proposal_sigma(pts.slice(0, i, end)));
}
return weights_from_sigma(torch::cat(chunks, 0).view(z_vals.sizes()), z_vals);
}
RenderOutput NeRFRenderer::volume_render(const torch::Tensor &rays_o,
const torch::Tensor &rays_d,
const torch::Tensor &z_vals,
int batch_size) const {
auto pts = rays_o.unsqueeze(-2) + rays_d.unsqueeze(-2) * z_vals.unsqueeze(-1);
auto view_dirs = rays_d / (rays_d.norm(2, -1, true) + 1e-8);
auto pts_flat = pts.reshape({-1, 3});
auto dirs_flat = view_dirs.unsqueeze(-2).expand(pts.sizes()).reshape({-1, 3});
std::vector<torch::Tensor> rgb_chunks, sigma_chunks;
for (int64_t i = 0; i < pts_flat.size(0); i += batch_size) {
const auto end = std::min<int64_t>(i + batch_size, pts_flat.size(0));
auto chunk = model_.forward(pts_flat.slice(0, i, end), dirs_flat.slice(0, i, end));
rgb_chunks.push_back(chunk.rgb);
sigma_chunks.push_back(chunk.sigma);
}
const int64_t N = z_vals.size(0), S = z_vals.size(1);
auto rgb = torch::cat(rgb_chunks, 0).view({N, S, 3});
auto sigma = torch::cat(sigma_chunks, 0).view({N, S});
auto weights = weights_from_sigma(sigma, z_vals);
RenderOutput out;
out.rgb = (weights.unsqueeze(-1) * rgb).sum(-2) +
(1.0 - weights.sum(-1, true)) * bg_color_;
out.depth = (weights * z_vals).sum(-1);
out.weights = weights;
return out;
}
torch::Tensor sample_pdf(const torch::Tensor &bins, const torch::Tensor &weights,
int n_samples, bool deterministic) {
const int64_t R = weights.size(0), M = weights.size(1);
// Normalise to a PDF, then a CDF with a leading zero to match the M + 1 edges.
auto w = weights + 1e-5f;
auto cdf = torch::cumsum(w / w.sum(-1, true), -1);
cdf = torch::cat({torch::zeros({R, 1}, cdf.options()), cdf}, -1);
auto u = deterministic ? torch::linspace(0.0f, 1.0f, n_samples, cdf.options())
.expand({R, n_samples})
.contiguous()
: torch::rand({R, n_samples}, cdf.options());
// Locate the CDF bin each u falls in, then invert it linearly.
auto inds = torch::searchsorted(cdf, u, /*out_int32=*/false, /*right=*/true);
auto pair = torch::stack({(inds - 1).clamp_min(0), inds.clamp_max(M)}, -1);
auto gather_pair = [&](const torch::Tensor &src) {
return torch::gather(src.unsqueeze(1).expand({R, n_samples, M + 1}), 2, pair);
};
auto cdf_g = gather_pair(cdf);
auto bins_g = gather_pair(bins);
auto denom = cdf_g.index({"...", 1}) - cdf_g.index({"...", 0});
denom = torch::where(denom < 1e-5f, torch::ones_like(denom), denom);
auto t = (u - cdf_g.index({"...", 0})) / denom;
return bins_g.index({"...", 0}) +
t * (bins_g.index({"...", 1}) - bins_g.index({"...", 0}));
}
namespace {
// Per-ray linear interpolation of fp (defined at sorted xp) at query points.
torch::Tensor interp_per_ray(const torch::Tensor &query, const torch::Tensor &xp,
const torch::Tensor &fp) {
auto inds = torch::searchsorted(xp, query, /*out_int32=*/false, /*right=*/true)
.clamp(1, xp.size(1) - 1);
auto lo = inds - 1;
auto x0 = torch::gather(xp, 1, lo), x1 = torch::gather(xp, 1, inds);
auto y0 = torch::gather(fp, 1, lo), y1 = torch::gather(fp, 1, inds);
return y0 + ((query - x0) / (x1 - x0 + 1e-7f)).clamp(0.0f, 1.0f) * (y1 - y0);
}
// Sample positions [R, S] -> S + 1 bin edges.
torch::Tensor to_edges(const torch::Tensor &z) {
return torch::cat({head(z), midpoints(z), z.index({"...", Slice(-1, None)})}, -1);
}
} // namespace
torch::Tensor interlevel_loss(const torch::Tensor &z_main, const torch::Tensor &w_main,
const torch::Tensor &z_prop, const torch::Tensor &w_prop) {
// Compare normalised distributions; the main one is a fixed target.
auto wm = (w_main / (w_main.sum(-1, true) + 1e-6)).detach();
auto wp = w_prop / (w_prop.sum(-1, true) + 1e-6);
// Proposal mass over each main bin is the difference of its interpolated CDF at
// the bin edges. Penalise where the proposal underestimates it.
auto cwp = torch::cat({torch::zeros({wp.size(0), 1}, wp.options()),
torch::cumsum(wp, -1)},
-1);
auto cw_at = interp_per_ray(to_edges(z_main), to_edges(z_prop), cwp);
auto excess = torch::relu(wm - (drop_first(cw_at) - drop_last(cw_at)));
return (excess * excess / (wm + 1e-6)).sum(-1).mean();
}