| 266 | } |
| 267 | |
| 268 | std::tuple<torch::Tensor, torch::Tensor> lambert_bwd(torch::Tensor nrm, torch::Tensor wi, torch::Tensor grad) |
| 269 | { |
| 270 | cudaStream_t stream = at::cuda::getCurrentCUDAStream(); |
| 271 | |
| 272 | // Extract input parameters. |
| 273 | LambertKernelParams p; |
| 274 | update_grid(p.gridSize, nrm, wi); |
| 275 | |
| 276 | // Choose launch parameters. |
| 277 | dim3 blockSize = getLaunchBlockSize(BLOCK_X, BLOCK_Y, p.gridSize); |
| 278 | dim3 gridSize = getLaunchGridSize(blockSize, p.gridSize); |
| 279 | |
| 280 | torch::Tensor nrm_grad, wi_grad; |
| 281 | p.nrm = make_cuda_tensor(nrm, p.gridSize, &nrm_grad); |
| 282 | p.wi = make_cuda_tensor(wi, p.gridSize, &wi_grad); |
| 283 | p.out = make_cuda_tensor(grad, p.gridSize); |
| 284 | |
| 285 | // Launch CUDA kernel. |
| 286 | void* args[] = { &p }; |
| 287 | NVDR_CHECK_CUDA_ERROR(cudaLaunchKernel((const void*)LambertBwdKernel, gridSize, blockSize, args, 0, stream)); |
| 288 | |
| 289 | return std::tuple<torch::Tensor, torch::Tensor>(nrm_grad, wi_grad); |
| 290 | } |
| 291 | |
| 292 | //------------------------------------------------------------------------ |
| 293 | // frostbite diffuse |
nothing calls this directly
no test coverage detected