pathtracer-diff / local_bind.cpp
phanerozoic's picture
v3: geometry gradients (dual-number interior + shadow and camera silhouette edge sampling)
1f9a369 verified
Raw
History Blame
9.13 kB
// Local development binding (torch.utils.cpp_extension JIT; Windows/MSVC).
// Mirrors the torch-ext binding's argument order exactly.
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include "geometry_grad.h"
#include "pathtracer_launch.h"
namespace {
const float* fptr(const torch::Tensor& t) {
return t.numel() ? t.const_data_ptr<float>() : nullptr;
}
const int* iptr(const torch::Tensor& t) {
return t.numel() ? t.const_data_ptr<int>() : nullptr;
}
PtdSceneArgs pack_args(
const torch::Tensor& tris, const torch::Tensor& mat_ids,
const torch::Tensor& uvs, const torch::Tensor& nodes_f,
const torch::Tensor& nodes_i, const torch::Tensor& light_faces,
const torch::Tensor& light_cdf, double total_light_area,
const torch::Tensor& tex, const torch::Tensor& tex_hdr,
const torch::Tensor& emi_tex, const torch::Tensor& emi_hdr,
const torch::Tensor& mat_type, const torch::Tensor& mat_rough,
const torch::Tensor& mat_ior, const torch::Tensor& med_sa,
const torch::Tensor& med_ss, double med_sbar, const torch::Tensor& env,
const torch::Tensor& env_cdf_m, const torch::Tensor& env_cdf_c,
const torch::Tensor& env_pdf, int64_t env_w, int64_t env_h,
int64_t max_bounces) {
TORCH_CHECK(tris.is_cuda() && tris.is_contiguous(), "tris");
TORCH_CHECK(uvs.numel() == tris.size(0) * 6, "uvs [F, 3, 2]");
TORCH_CHECK(tex_hdr.size(0) <= 64, "at most 64 materials");
TORCH_CHECK(max_bounces >= 1 && max_bounces <= 16, "bounces in [1,16]");
PtdSceneArgs a;
a.tris = tris.const_data_ptr<float>();
a.mat_ids = mat_ids.const_data_ptr<int>();
a.uvs = uvs.const_data_ptr<float>();
a.n_faces = (int)tris.size(0);
a.nodes_f = nodes_f.const_data_ptr<float>();
a.nodes_i = nodes_i.const_data_ptr<int>();
a.n_nodes = (int)nodes_f.size(0);
a.light_faces = iptr(light_faces);
a.light_cdf = fptr(light_cdf);
a.n_lights = (int)light_faces.numel();
a.total_light_area = (float)total_light_area;
a.tex = tex.const_data_ptr<float>();
a.tex_hdr = tex_hdr.const_data_ptr<int>();
a.n_texels = (int)tex.size(0);
a.emi_tex = emi_tex.const_data_ptr<float>();
a.emi_hdr = emi_hdr.const_data_ptr<int>();
a.n_emi_texels = (int)emi_tex.size(0);
a.mat_type = mat_type.const_data_ptr<int>();
a.mat_rough = mat_rough.const_data_ptr<float>();
a.mat_ior = mat_ior.const_data_ptr<float>();
a.n_mats = (int)tex_hdr.size(0);
a.med_sa = fptr(med_sa);
a.med_ss = fptr(med_ss);
a.med_sbar = (float)med_sbar;
a.has_med = med_sa.numel() ? 1 : 0;
a.env = fptr(env);
a.env_w = (int)env_w;
a.env_h = (int)env_h;
a.env_cdf_m = fptr(env_cdf_m);
a.env_cdf_c = fptr(env_cdf_c);
a.env_pdf = fptr(env_pdf);
return a;
}
void pt_forward(torch::Tensor tris, torch::Tensor mat_ids, torch::Tensor uvs,
torch::Tensor nodes_f, torch::Tensor nodes_i,
torch::Tensor light_faces, torch::Tensor light_cdf,
double total_light_area, torch::Tensor tex,
torch::Tensor tex_hdr, torch::Tensor emi_tex,
torch::Tensor emi_hdr, torch::Tensor mat_type,
torch::Tensor mat_rough, torch::Tensor mat_ior,
torch::Tensor med_sa, torch::Tensor med_ss, double med_sbar,
torch::Tensor env, torch::Tensor env_cdf_m,
torch::Tensor env_cdf_c, torch::Tensor env_pdf, int64_t env_w,
int64_t env_h, torch::Tensor cam, int64_t spp,
int64_t max_bounces, int64_t mode, int64_t seed,
torch::Tensor image) {
PtdSceneArgs a = pack_args(tris, mat_ids, uvs, nodes_f, nodes_i, light_faces,
light_cdf, total_light_area, tex, tex_hdr,
emi_tex, emi_hdr, mat_type, mat_rough, mat_ior,
med_sa, med_ss, med_sbar, env, env_cdf_m,
env_cdf_c, env_pdf, env_w, env_h, max_bounces);
const at::cuda::CUDAGuard guard(tris.device());
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
torch::Tensor cam_h = cam.to(torch::kFloat32).to(torch::kCPU).contiguous();
ptd_forward_launch(&a, cam_h.const_data_ptr<float>(), (int)image.size(0),
(int)image.size(1), (int)spp, (int)max_bounces,
(int)mode, (long long)seed, image.data_ptr<float>(),
stream);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void pt_backward(torch::Tensor tris, torch::Tensor mat_ids, torch::Tensor uvs,
torch::Tensor nodes_f, torch::Tensor nodes_i,
torch::Tensor light_faces, torch::Tensor light_cdf,
double total_light_area, torch::Tensor tex,
torch::Tensor tex_hdr, torch::Tensor emi_tex,
torch::Tensor emi_hdr, torch::Tensor mat_type,
torch::Tensor mat_rough, torch::Tensor mat_ior,
torch::Tensor med_sa, torch::Tensor med_ss, double med_sbar,
torch::Tensor env, torch::Tensor env_cdf_m,
torch::Tensor env_cdf_c, torch::Tensor env_pdf, int64_t env_w,
int64_t env_h, torch::Tensor cam, int64_t spp,
int64_t max_bounces, int64_t mode, int64_t seed,
torch::Tensor grad_image, torch::Tensor grad_tex,
torch::Tensor grad_emi_tex, torch::Tensor grad_env,
torch::Tensor grad_med) {
PtdSceneArgs a = pack_args(tris, mat_ids, uvs, nodes_f, nodes_i, light_faces,
light_cdf, total_light_area, tex, tex_hdr,
emi_tex, emi_hdr, mat_type, mat_rough, mat_ior,
med_sa, med_ss, med_sbar, env, env_cdf_m,
env_cdf_c, env_pdf, env_w, env_h, max_bounces);
const at::cuda::CUDAGuard guard(tris.device());
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
torch::Tensor cam_h = cam.to(torch::kFloat32).to(torch::kCPU).contiguous();
ptd_backward_launch(&a, cam_h.const_data_ptr<float>(),
(int)grad_image.size(0), (int)grad_image.size(1),
(int)spp, (int)max_bounces, (int)mode, (long long)seed,
grad_image.const_data_ptr<float>(),
grad_tex.data_ptr<float>(),
grad_emi_tex.data_ptr<float>(),
grad_env.numel() ? grad_env.data_ptr<float>() : nullptr,
grad_med.numel() ? grad_med.data_ptr<float>() : nullptr,
stream);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void pt_geometry_grad(torch::Tensor tris, torch::Tensor mat_ids,
torch::Tensor uvs, torch::Tensor nodes_f,
torch::Tensor nodes_i, torch::Tensor light_faces,
torch::Tensor light_cdf, double total_light_area,
torch::Tensor tex, torch::Tensor tex_hdr,
torch::Tensor emi_tex, torch::Tensor emi_hdr,
torch::Tensor mat_type, torch::Tensor mat_rough,
torch::Tensor mat_ior, torch::Tensor face_verts,
torch::Tensor edges, torch::Tensor edge_cdf,
torch::Tensor cam, int64_t spp, int64_t edge_samples,
int64_t seed, torch::Tensor grad_image,
torch::Tensor grad_verts) {
torch::Tensor z = torch::zeros(0, tris.options());
PtdSceneArgs a = pack_args(tris, mat_ids, uvs, nodes_f, nodes_i,
light_faces, light_cdf, total_light_area, tex,
tex_hdr, emi_tex, emi_hdr, mat_type, mat_rough,
mat_ior, z, z, 0.0, z, z, z, z, 0, 0, 4);
const at::cuda::CUDAGuard guard(tris.device());
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
torch::Tensor cam_h = cam.to(torch::kFloat32).to(torch::kCPU).contiguous();
int H = (int)grad_image.size(0), W = (int)grad_image.size(1);
ptd_geo_interior_launch(&a, face_verts.const_data_ptr<int>(),
cam_h.const_data_ptr<float>(), H, W, (int)spp,
(long long)seed,
grad_image.const_data_ptr<float>(),
grad_verts.data_ptr<float>(), stream);
if (edges.numel())
ptd_geo_boundary_launch(&a, face_verts.const_data_ptr<int>(),
edges.const_data_ptr<int>(),
(int)edges.size(0),
edge_cdf.const_data_ptr<float>(),
cam_h.const_data_ptr<float>(), H, W,
(int)edge_samples, (long long)seed,
grad_image.const_data_ptr<float>(),
grad_verts.data_ptr<float>(), stream);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("pt_forward", &pt_forward, "pathtracer-diff forward");
m.def("pt_backward", &pt_backward, "pathtracer-diff backward");
m.def("pt_geometry_grad", &pt_geometry_grad,
"pathtracer-diff geometry gradients (direct lighting)");
}