File size: 2,035 Bytes
b6045d5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
/*

 * Copyright (C) 2023, Inria

 * GRAPHDECO research group, https://team.inria.fr/graphdeco

 * All rights reserved.

 *

 * This software is free for non-commercial, research and evaluation use 

 * under the terms of the LICENSE.md file.

 *

 * For inquiries contact  george.drettakis@inria.fr

 */

#pragma once
#include <torch/extension.h>
#include <cstdio>
#include <tuple>
#include <string>
	
std::tuple<int, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>
RasterizeGaussiansCUDA(

	const torch::Tensor& background,

	const torch::Tensor& means3D,

    const torch::Tensor& colors,

    const torch::Tensor& opacity,

	const torch::Tensor& scales,

	const torch::Tensor& rotations,

	const float scale_modifier,

	const torch::Tensor& cov3D_precomp,

	const torch::Tensor& viewmatrix,

	const torch::Tensor& projmatrix,

	const float tan_fovx, 

	const float tan_fovy,

    const int image_height,

    const int image_width,

	const torch::Tensor& sh,

	const int degree,

	const torch::Tensor& campos,

	const bool prefiltered,

	const bool debug);

std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>
 RasterizeGaussiansBackwardCUDA(

 	const torch::Tensor& background,

	const torch::Tensor& means3D,

	const torch::Tensor& radii,

    const torch::Tensor& colors,

	const torch::Tensor& scales,

	const torch::Tensor& rotations,

	const float scale_modifier,

	const torch::Tensor& cov3D_precomp,

	const torch::Tensor& viewmatrix,

    const torch::Tensor& projmatrix,

	const float tan_fovx, 

	const float tan_fovy,

    const torch::Tensor& dL_dout_color,

	const torch::Tensor& sh,

	const int degree,

	const torch::Tensor& campos,

	const torch::Tensor& geomBuffer,

	const int R,

	const torch::Tensor& binningBuffer,

	const torch::Tensor& imageBuffer,

	const bool debug);
		
torch::Tensor markVisible(

		torch::Tensor& means3D,

		torch::Tensor& viewmatrix,

		torch::Tensor& projmatrix);