// Local development binding (torch.utils.cpp_extension JIT; Windows/MSVC). // Mirrors the torch-ext binding's argument order exactly. #include #include #include #include "geometry_grad.h" #include "pathtracer_launch.h" namespace { const float* fptr(const torch::Tensor& t) { return t.numel() ? t.const_data_ptr() : nullptr; } const int* iptr(const torch::Tensor& t) { return t.numel() ? t.const_data_ptr() : 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(); a.mat_ids = mat_ids.const_data_ptr(); a.uvs = uvs.const_data_ptr(); a.n_faces = (int)tris.size(0); a.nodes_f = nodes_f.const_data_ptr(); a.nodes_i = nodes_i.const_data_ptr(); 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(); a.tex_hdr = tex_hdr.const_data_ptr(); a.n_texels = (int)tex.size(0); a.emi_tex = emi_tex.const_data_ptr(); a.emi_hdr = emi_hdr.const_data_ptr(); a.n_emi_texels = (int)emi_tex.size(0); a.mat_type = mat_type.const_data_ptr(); a.mat_rough = mat_rough.const_data_ptr(); a.mat_ior = mat_ior.const_data_ptr(); 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(), (int)image.size(0), (int)image.size(1), (int)spp, (int)max_bounces, (int)mode, (long long)seed, image.data_ptr(), 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(), (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(), grad_tex.data_ptr(), grad_emi_tex.data_ptr(), grad_env.numel() ? grad_env.data_ptr() : nullptr, grad_med.numel() ? grad_med.data_ptr() : 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(), cam_h.const_data_ptr(), H, W, (int)spp, (long long)seed, grad_image.const_data_ptr(), grad_verts.data_ptr(), stream); if (edges.numel()) ptd_geo_boundary_launch(&a, face_verts.const_data_ptr(), edges.const_data_ptr(), (int)edges.size(0), edge_cdf.const_data_ptr(), cam_h.const_data_ptr(), H, W, (int)edge_samples, (long long)seed, grad_image.const_data_ptr(), grad_verts.data_ptr(), 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)"); }