move kernel instantiation to separate files (wip)

This commit is contained in:
Harry Chen
2026-03-09 22:30:02 -04:00
parent e2929eab54
commit 5cc5f3330c
142 changed files with 7644 additions and 2704 deletions
+11 -3
View File
@@ -77,7 +77,9 @@ def get_extensions():
from torch.utils.cpp_extension import CUDAExtension
extensions_dir = os.path.abspath(os.path.join("spirulae_splat", "splat", "cuda", "csrc"))
sources = glob.glob(os.path.join(extensions_dir, "*.cu")) + \
sources = \
glob.glob(os.path.join(os.path.join(extensions_dir, "generated", "kernel_instantiation"), "*.cu")) + \
glob.glob(os.path.join(extensions_dir, "*.cu")) + \
glob.glob(os.path.join(extensions_dir, "*.cpp"))
sources = [path for path in sources if "hip" not in path]
@@ -114,14 +116,15 @@ def get_extensions():
nvcc_flags = os.getenv("NVCC_FLAGS", "")
nvcc_flags = [] if nvcc_flags == "" else nvcc_flags.split(" ")
nvcc_flags += ["-O3", "--use_fast_math"]
# nvcc_flags += ["-rdc=true", "-dlto"]
# nvcc_flags += ["--extra-device-vectorization"]
if LINE_INFO:
nvcc_flags += ["-lineinfo", "--generate-line-info", "--source-in-ptx"]
nvcc_flags += [
# "-Xptxas", "-v",
"-Xptxas", "--warn-on-double-precision-use",
"-Xptxas", "--warn-on-local-memory-usage",
"-Xptxas", "--warn-on-spills"
# "-Xptxas", "--warn-on-local-memory-usage",
# "-Xptxas", "--warn-on-spills"
]
if torch.version.hip:
# USE_ROCM was added to later versions of PyTorch.
@@ -159,6 +162,7 @@ def get_extensions():
undef_macros=undef_macros,
extra_compile_args=extra_compile_args,
extra_link_args=extra_link_args,
# dlink=True
)
return [extension]
@@ -179,9 +183,13 @@ for filename in [
'PerPixelLoss',
'PixelWise',
'Projection',
'ProjectionFwd',
'ProjectionBwd',
'ProjectionHeteroFwd',
'ProjectionHeteroBwd',
'Rasterization',
'RasterizationEval3DFwd',
'RasterizationEval3DBwd',
'RasterizationSortedEval3DFwd',
'RasterizationSortedEval3DBwd',
'Optimizer',
+624
View File
@@ -0,0 +1,624 @@
#!/usr/bin/env python3
"""
cuda_resource_viewer.py — Pretty-print CUDA kernel resource usage from compiled binaries.
Usage:
python cuda_resource_viewer.py <path/to/binary.so>
python cuda_resource_viewer.py <path/to/binary.so> --no-color
python cuda_resource_viewer.py <path/to/binary.so> --sort-by reg|shared|const|name
"""
import sys
import re
import subprocess
import shutil
import argparse
from dataclasses import dataclass, field
from typing import Optional
# ---------------------------------------------------------------------------
# ANSI helpers
# ---------------------------------------------------------------------------
NO_COLOR = False # toggled by --no-color
def _c(*codes: int) -> str:
if NO_COLOR:
return ""
return f"\033[{';'.join(map(str, codes))}m"
RESET = lambda: _c(0)
BOLD = lambda: _c(1)
DIM = lambda: _c(2)
ITALIC = lambda: _c(3)
UNDERLINE = lambda: _c(4)
# Foreground colours
BLACK = lambda: _c(30)
RED = lambda: _c(31)
GREEN = lambda: _c(32)
YELLOW = lambda: _c(33)
BLUE = lambda: _c(34)
MAGENTA = lambda: _c(35)
CYAN = lambda: _c(36)
WHITE = lambda: _c(37)
# Bright variants
BRIGHT_RED = lambda: _c(91)
BRIGHT_GREEN = lambda: _c(92)
BRIGHT_YELLOW = lambda: _c(93)
BRIGHT_BLUE = lambda: _c(94)
BRIGHT_MAGENTA = lambda: _c(95)
BRIGHT_CYAN = lambda: _c(96)
BRIGHT_WHITE = lambda: _c(97)
# Background colours
BG_RED = lambda: _c(41)
BG_YELLOW = lambda: _c(43)
BG_GREEN = lambda: _c(42)
BG_BLUE = lambda: _c(44)
def styled(text: str, *fns) -> str:
prefix = "".join(f() for f in fns)
if not prefix:
return text
return f"{prefix}{text}{RESET()}"
# ---------------------------------------------------------------------------
# Resource thresholds for colour-coding
# ---------------------------------------------------------------------------
REG_THRESHOLDS = [(32, BRIGHT_GREEN), (64, BRIGHT_YELLOW), (96, BRIGHT_RED), (999, RED)]
SHARED_THRESHOLDS = [(0, DIM), (16384, BRIGHT_GREEN), (32768, BRIGHT_YELLOW), (999999, BRIGHT_RED)]
def color_value(value: int, thresholds: list) -> str:
for limit, color_fn in thresholds:
if value <= limit:
return styled(str(value), color_fn)
return str(value)
def reg_color(v: int) -> str:
return color_value(v, REG_THRESHOLDS)
def shared_color(v: int) -> str:
if v == 0:
return styled("0", DIM)
for limit, color_fn in SHARED_THRESHOLDS:
if v <= limit:
return styled(str(v), color_fn)
return str(v)
def const_color(v: int) -> str:
if v == 0:
return styled("0", DIM)
if v < 512:
return styled(str(v), BRIGHT_GREEN)
if v < 1024:
return styled(str(v), BRIGHT_YELLOW)
return styled(str(v), BRIGHT_RED)
def generic_color(v: int) -> str:
if v == 0:
return styled("0", DIM)
return styled(str(v), BRIGHT_CYAN)
# ---------------------------------------------------------------------------
# Data model
# ---------------------------------------------------------------------------
@dataclass
class KernelInfo:
mangled_name: str
demangled_name: Optional[str] = None
reg: int = 0
stack: int = 0
shared: int = 0
local: int = 0
constants: dict = field(default_factory=dict) # {idx: bytes}
texture: int = 0
surface: int = 0
sampler: int = 0
@property
def display_name(self) -> str:
return self.demangled_name or self.mangled_name
@dataclass
class FatbinSection:
arch: str = ""
code_version: str = ""
host: str = ""
compile_size: str = ""
identifier: str = ""
common_global: int = 0
common_constants: dict = field(default_factory=dict) # {idx: bytes}
kernels: list = field(default_factory=list)
# ---------------------------------------------------------------------------
# Parsing
# ---------------------------------------------------------------------------
RES_PATTERN = re.compile(
r"REG:(\d+)\s+STACK:(\d+)\s+SHARED:(\d+)\s+LOCAL:(\d+)"
r"((?:\s+CONSTANT\[\d+\]:\d+)+)"
r"\s+TEXTURE:(\d+)\s+SURFACE:(\d+)\s+SAMPLER:(\d+)"
)
CONST_PATTERN = re.compile(r"CONSTANT\[(\d+)\]:(\d+)")
COMMON_GLOBAL_PATTERN = re.compile(r"GLOBAL:(\d+)")
COMMON_CONST_PATTERN = re.compile(r"CONSTANT\[(\d+)\]:(\d+)")
def parse_output(text: str) -> list[FatbinSection]:
sections: list[FatbinSection] = []
current: Optional[FatbinSection] = None
in_resource_block = False
pending_func: Optional[str] = None
for line in text.splitlines():
stripped = line.strip()
# Detect new ELF fatbin section
if stripped.startswith("Fatbin elf code:"):
current = FatbinSection()
sections.append(current)
in_resource_block = False
pending_func = None
continue
# Skip PTX sections entirely
if stripped.startswith("Fatbin ptx code:"):
current = None
continue
if current is None:
continue
# Header fields
if stripped.startswith("arch ="):
current.arch = stripped.split("=", 1)[1].strip()
elif stripped.startswith("code version ="):
current.code_version = stripped.split("=", 1)[1].strip()
elif stripped.startswith("host ="):
current.host = stripped.split("=", 1)[1].strip()
elif stripped.startswith("compile_size ="):
current.compile_size = stripped.split("=", 1)[1].strip()
elif stripped.startswith("identifier ="):
current.identifier = stripped.split("=", 1)[1].strip()
elif stripped == "Resource usage:":
in_resource_block = True
elif not in_resource_block:
continue
elif stripped == "Common:":
pass
elif stripped.startswith("GLOBAL:"):
m = COMMON_GLOBAL_PATTERN.search(stripped)
if m:
current.common_global = int(m.group(1))
for cm in COMMON_CONST_PATTERN.finditer(stripped):
current.common_constants[int(cm.group(1))] = int(cm.group(2))
elif stripped.startswith("Function "):
pending_func = stripped[len("Function "):]
if pending_func.endswith(":"):
pending_func = pending_func[:-1]
elif pending_func and (stripped.startswith("REG:") or "REG:" in stripped):
m = RES_PATTERN.search(stripped)
if m:
consts = {int(cm.group(1)): int(cm.group(2))
for cm in CONST_PATTERN.finditer(m.group(5))}
k = KernelInfo(
mangled_name=pending_func,
reg=int(m.group(1)),
stack=int(m.group(2)),
shared=int(m.group(3)),
local=int(m.group(4)),
constants=consts,
texture=int(m.group(6)),
surface=int(m.group(7)),
sampler=int(m.group(8)),
)
current.kernels.append(k)
pending_func = None
return [s for s in sections if s.identifier or s.kernels]
# ---------------------------------------------------------------------------
# C++ demangling
# ---------------------------------------------------------------------------
def demangle_names(kernels: list[KernelInfo]) -> None:
"""Attempt to demangle mangled C++ names via c++filt."""
if not kernels:
return
c_filt = shutil.which("c++filt")
if not c_filt:
return
names = [k.mangled_name for k in kernels]
try:
result = subprocess.run(
[c_filt] + names,
capture_output=True, text=True, timeout=10
)
demangled = result.stdout.strip().splitlines()
for k, d in zip(kernels, demangled):
if d and d != k.mangled_name:
k.demangled_name = d
except Exception:
pass
def demangle_all(sections: list[FatbinSection]) -> None:
all_kernels = [k for s in sections for k in s.kernels]
demangle_names(all_kernels)
# ---------------------------------------------------------------------------
# Display helpers
# ---------------------------------------------------------------------------
BOX_H = "─"
BOX_V = "│"
BOX_TL = "╭"
BOX_TR = "╮"
BOX_BL = "╰"
BOX_BR = "╯"
BOX_LM = "├"
BOX_RM = "┤"
TERM_WIDTH = 100
def hr(char: str = BOX_H, width: int = TERM_WIDTH, color=None) -> str:
line = char * width
return styled(line, color) if color else line
def box_line(text: str, width: int = TERM_WIDTH) -> str:
inner = width - 4
return f" {styled(BOX_V, CYAN)} {text:<{inner}}{styled(BOX_V, CYAN)}"
def section_header(title: str) -> str:
pad = TERM_WIDTH - len(title) + 6
left = pad // 2
right = pad - left
return (
" " +
styled(BOX_TL + BOX_H * (left + 1), CYAN) +
styled(f" {title} ", BOLD, BRIGHT_WHITE) +
styled(BOX_H * (right + 1) + BOX_TR, CYAN)
)
def section_footer() -> str:
return styled(" " + BOX_BL + BOX_H * (TERM_WIDTH - 3) + BOX_BR, CYAN)
def format_kernel_resources(k: KernelInfo) -> list[str]:
"""Return lines describing kernel resources (non-zero highlighted)."""
items = []
line_length = 0
# REG is always shown
items.append(f"{styled('REG', BOLD)}:{reg_color(k.reg)}")
line_length += len('REG:') + len(str(k.reg))
# STACK – show only if nonzero
if k.stack:
items.append(f"{styled('STACK', BOLD)}:{styled(str(k.stack), BRIGHT_RED)}")
line_length += len('STACK:') + len(str(k.stack))
# SHARED – always show (important for occupancy)
if k.shared:
items.append(f"{styled('SHARED', BOLD)}:{shared_color(k.shared)}")
line_length += len('SHARED:') + len(str(k.shared))
else:
items.append(f"{styled('SHARED', DIM)}:{styled('0', DIM)}")
line_length += len('SHARED:0')
# LOCAL
if k.local:
items.append(f"{styled('LOCAL', BOLD)}:{BRIGHT_RED()}{k.local}{RESET()}")
line_length += len('LOCAL:') + len(str(k.local))
# CONSTANTS
for idx in sorted(k.constants):
v = k.constants[idx]
label = styled(f"CONST[{idx}]", DIM)
val = const_color(v)
items.append(f"{label}:{val}")
line_length += len(f"CONST[{idx}]:") + len(str(val))
# TEXTURE / SURFACE / SAMPLER (only if nonzero)
for label, val in [("TEX", k.texture), ("SURF", k.surface), ("SAMP", k.sampler)]:
if val:
items.append(f"{styled(label, DIM)}:{styled(str(val), BRIGHT_MAGENTA)}")
line_length += len(f"{label}:") + len(str(val))
res_line = " ".join(items)
line_length += len(items) - 1
while line_length < 60:
res_line += " "
line_length += 1
return res_line, line_length
def reg_bar(value: int, max_val: int = 255, width: int = 20) -> str:
"""Visual register usage bar."""
if max_val == 0:
return ""
filled = round(value / max_val * width)
filled = max(0, min(width, filled))
ratio = value / max_val
if ratio < 0.5:
bar_color = BRIGHT_GREEN
elif ratio < 0.75:
bar_color = BRIGHT_YELLOW
else:
bar_color = BRIGHT_RED
bar = styled("█" * filled, bar_color) + styled("░" * (width - filled), DIM)
pct = f"{ratio*100:4.0f}%"
return f"[{bar}] {styled(pct, DIM)}"
def truncate_name(name: str, max_len: int = TERM_WIDTH - 8) -> str:
if len(name) <= max_len:
return name + ' '*(max_len-len(name))
if '>(' in name:
name = name[:max(name.rfind('>(')+1, max_len-1)]
else:
name = name[:max_len - 1]
return name + "…"
def print_section(section: FatbinSection, sort_key: Optional[str] = None, idx: int = 0) -> None:
kernels = section.kernels
if sort_key:
if sort_key == "reg":
kernels = sorted(kernels, key=lambda k: -k.reg)
elif sort_key == "shared":
kernels = sorted(kernels, key=lambda k: -k.shared)
elif sort_key == "const":
kernels = sorted(kernels, key=lambda k: -max(k.constants.values(), default=0))
elif sort_key == "name":
kernels = sorted(kernels, key=lambda k: k.display_name)
# ---- Section header ----
src_name = section.identifier.split("/")[-1] if "/" in section.identifier else section.identifier
title = f" {styled(src_name or f'Section {idx+1}', BOLD, BRIGHT_WHITE)} "
print()
print(section_header(title.strip()))
# ---- Metadata row ----
meta_parts = [
f"{styled('arch', DIM)}={styled(section.arch, BRIGHT_CYAN)}",
f"{styled('ver', DIM)}={styled(section.code_version, BRIGHT_CYAN)}",
f"{styled('host', DIM)}={styled(section.host, DIM)}",
f"{styled('bits', DIM)}={styled(section.compile_size.replace('bit',''), DIM)}",
f"{styled('kernels', DIM)}={styled(str(len(kernels)), BRIGHT_WHITE, BOLD)}",
]
print(f" {styled('│', CYAN)} " + " ".join(meta_parts))
# ---- Full identifier path ----
if section.identifier:
print(f" {styled('│', CYAN)} {styled('src', DIM)}: {styled(section.identifier, DIM)}")
# ---- Common resources ----
common_parts = []
if section.common_global:
common_parts.append(f"{styled('GLOBAL', DIM)}:{generic_color(section.common_global)}")
for idx2 in sorted(section.common_constants):
v = section.common_constants[idx2]
common_parts.append(f"{styled(f'CONST[{idx2}]', DIM)}:{const_color(v)}")
if common_parts:
print(f" {styled('│', CYAN)} {styled('common', ITALIC, DIM)}: {' '.join(common_parts)}")
# ---- Separator ----
print(f" {styled(BOX_LM + BOX_H * (TERM_WIDTH - 3) + BOX_RM, CYAN)}")
if not kernels:
print(f" {styled('│', CYAN)} {styled('(no kernels)', DIM)}")
else:
max_reg = max((k.reg for k in kernels), default=128)
for i, k in enumerate(sorted(kernels, key=lambda k: k.display_name)):
# Kernel name
raw = k.display_name
# Try to highlight template params dimly
name_display = format_display_name(raw)
print(f" {styled('│', CYAN)} {styled(f'❯ {name_display}', BRIGHT_WHITE, BOLD)} {styled('│', CYAN)}")
# Resources row
res_line, line_length = format_kernel_resources(k)
# bar = reg_bar(k.reg, max(max_reg, 32))
bar = reg_bar(k.reg, 128)
whitespace = ' ' * (TERM_WIDTH - line_length - 32)
print(f" {styled('│', CYAN)} {res_line} {bar} {whitespace} {styled('│', CYAN)}")
# if i < len(kernels) - 1:
# print(f" {styled('│', CYAN)} {styled('·' * (TERM_WIDTH - 6), DIM)} {styled('│', CYAN)}")
print(section_footer())
def format_display_name(name: str) -> str:
"""Highlight template params and namespaces with dim styling."""
name = truncate_name(name)
# Color angle brackets for templates
if NO_COLOR:
return name
result = ""
depth = 0
for ch in name:
if ch == "<":
depth += 1
result += BRIGHT_MAGENTA() + ch + (DIM() if depth > 0 else "")
elif ch == ">":
depth -= 1
result += (BRIGHT_MAGENTA() if depth == 0 else "") + ch + (RESET() + BRIGHT_WHITE() if depth == 0 else "")
elif ch == ":" and depth == 0:
result += DIM() + ch
else:
if depth > 0:
result += DIM() + ch
else:
result += ch
result += RESET()
return result
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def run_cuobjdump(binary_path: str) -> str:
tool = shutil.which("cuobjdump")
if not tool:
print(styled("✗ cuobjdump not found in PATH. Is CUDA toolkit installed?", BRIGHT_RED, BOLD),
file=sys.stderr)
sys.exit(1)
try:
result = subprocess.run(
[tool, "--dump-resource-usage", binary_path],
capture_output=True, text=True, timeout=60
)
if result.returncode != 0:
print(styled(f"✗ cuobjdump exited with code {result.returncode}:", BRIGHT_RED),
file=sys.stderr)
print(result.stderr, file=sys.stderr)
sys.exit(result.returncode)
return result.stdout
except FileNotFoundError:
print(styled("✗ cuobjdump not found.", BRIGHT_RED), file=sys.stderr)
sys.exit(1)
except subprocess.TimeoutExpired:
print(styled("✗ cuobjdump timed out.", BRIGHT_RED), file=sys.stderr)
sys.exit(1)
def print_summary(sections: list[FatbinSection]) -> None:
total_kernels = sum(len(s.kernels) for s in sections)
all_kernels = [k for s in sections for k in s.kernels]
if not all_kernels:
return
max_reg_k = max(all_kernels, key=lambda k: k.reg)
max_shared_k = max(all_kernels, key=lambda k: k.shared)
print()
print(styled(" SUMMARY ", BOLD, BRIGHT_WHITE, BG_BLUE) +
styled(f" {len(sections)} ELF section(s) {total_kernels} kernel(s) total", DIM))
print()
# Per-architecture table
archs: dict[str, list[KernelInfo]] = {}
for s in sections:
archs.setdefault(s.arch, []).extend(s.kernels)
header = (
styled(f" {'ARCH':<10}", BOLD) +
styled(f"{'KERNELS':>8}", BOLD) +
styled(f"{'MAX REG':>10}", BOLD) +
styled(f"{'AVG REG':>10}", BOLD) +
styled(f"{'MAX SHARED':>12}", BOLD)
)
print(header)
print(styled(" " + "─" * (len("ARCH") + 8 + 10 + 10 + 12 + 8), DIM))
for arch, ks in sorted(archs.items()):
avg_reg = sum(k.reg for k in ks) / len(ks) if ks else 0
mx_reg = max((k.reg for k in ks), default=0)
mx_sh = max((k.shared for k in ks), default=0)
print(
styled(f" {arch:<10}", BRIGHT_CYAN) +
f"{len(ks):>8}" +
f" {reg_color(mx_reg):>8}" + " " * 2 +
f" {styled(f'{avg_reg:.1f}', BRIGHT_WHITE):>8}" + " " * 2 +
f" {shared_color(mx_sh):>8}"
)
print()
print(styled(" Highest REG: ", DIM) +
styled(f"{max_reg_k.reg}", BOLD, BRIGHT_RED) +
" " + styled(truncate_name(max_reg_k.display_name, TERM_WIDTH-40), DIM))
if max_shared_k.shared > 0:
print(styled(" Highest SHARED:", DIM) +
styled(f" {max_shared_k.shared}", BOLD, BRIGHT_YELLOW) +
" " + styled(truncate_name(max_shared_k.display_name, TERM_WIDTH-40), DIM))
print()
def main() -> None:
global NO_COLOR
parser = argparse.ArgumentParser(
description="Pretty-print CUDA kernel resource usage from a compiled binary.",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog=__doc__,
)
parser.add_argument("binary", help="Path to CUDA binary (.so, .dll, …)")
parser.add_argument("--no-color", action="store_true", help="Disable ANSI colour output")
parser.add_argument(
"--sort-by",
choices=["reg", "shared", "const", "name"],
default=None,
help="Sort kernels within each section",
)
parser.add_argument(
"--raw",
action="store_true",
help="Print raw cuobjdump output and exit (for debugging)",
)
args = parser.parse_args()
NO_COLOR = args.no_color or not sys.stdout.isatty()
raw = run_cuobjdump(args.binary)
if args.raw:
print(raw)
return
sections = parse_output(raw)
demangle_all(sections)
if not sections:
print(styled("No ELF fatbin sections with resource usage found.", BRIGHT_YELLOW))
sys.exit(0)
# Banner
print()
print(styled("━" * TERM_WIDTH, CYAN))
print(styled(" CUDA KERNEL RESOURCE VIEWER", BOLD, BRIGHT_CYAN) +
styled(f" · {args.binary}", DIM))
print(styled("━" * TERM_WIDTH, CYAN))
print_summary(sections)
for i, section in enumerate(sections):
print_section(section, sort_key=args.sort_by, idx=i)
print()
if __name__ == "__main__":
main()
@@ -0,0 +1,218 @@
import os
from pathlib import Path
from typing import Tuple, List, Dict, Optional
import re
CSRC_DIR = Path("spirulae_splat", "splat", "cuda", "csrc")
DST_DIR = Path("spirulae_splat", "splat", "cuda", "csrc", "generated", "kernel_instantiation")
def extract_kernel_definition(header_src: Path, kernel_name: str):
header_src = CSRC_DIR / header_src
src = open(header_src, "r").read()
pattern = re.compile(
f"""(void\\s+{kernel_name}(.*?)\\)\\s*;)""",
re.MULTILINE | re.VERBOSE | re.DOTALL
)
matches = pattern.findall(src)
assert len(matches) == 1, matches
return matches[0][0]
def generate_kernel_instantiation(
predix: str,
template_definition: str,
map_header: List[Optional[str]],
maps: List[List[str]],
includes: List[List[str]]
):
assert len(maps) == len(includes)
for map, includes in zip(maps, includes):
definition = "template " + template_definition
assert len(map) == len(map_header)
for src, dst in zip(map_header, map):
if src is not None:
definition = definition.replace(src, dst)
template = "<\n " + ',\n '.join(map) + "\n>"
definition = definition.replace('(', template + "(", 1)
filename = [predix]
for name in map:
if '<' in name:
name = name.replace('<', '_').replace('>', '')
match = re.search(r'\w+$', name)
assert match, name
filename.append(match.group(0))
filename = '_'.join(filename) + ".cu"
with open(DST_DIR / filename, 'w') as fp:
fp.write("// This file is auto generated by `generate_kernel_instantiation.py`\n\n")
fp.write("#define NO_TORCH\n")
for include in includes:
fp.write(f"#include \"{include}\"\n")
fp.write("\n")
fp.write(definition)
fp.write("\n")
print("Generated", filename)
def generate_ProjectionFwd():
definition = extract_kernel_definition("ProjectionFwd.cu", "projection_fused_fwd_kernel_wrapper")
map_header = ["typename SplatPrimitive", None]
map_body = [
("Vanilla3DGS", "gsplat::CameraModelType::PINHOLE"),
("Vanilla3DGS", "gsplat::CameraModelType::FISHEYE"),
("MipSplatting", "gsplat::CameraModelType::PINHOLE"),
("MipSplatting", "gsplat::CameraModelType::FISHEYE"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<2>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<2>", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<3>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<3>", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<4>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<4>", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<5>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<5>", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<6>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<6>", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<7>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<7>", "gsplat::CameraModelType::FISHEYE"),
("SphericalVoronoi3DGUT<8>", "gsplat::CameraModelType::PINHOLE"),
("SphericalVoronoi3DGUT<8>", "gsplat::CameraModelType::FISHEYE"),
("OpaqueTriangle", "gsplat::CameraModelType::PINHOLE"),
("OpaqueTriangle", "gsplat::CameraModelType::FISHEYE"),
("VoxelPrimitive", "gsplat::CameraModelType::PINHOLE"),
("VoxelPrimitive", "gsplat::CameraModelType::FISHEYE"),
]
includes = [("Primitive3DGS.cuh", "ProjectionFwd_kernel.cuh")] * 4 + \
[("Primitive3DGUT.cuh", "ProjectionFwd_kernel.cuh")] * 2 + \
[("Primitive3DGUT_SV.cuh", "ProjectionFwd_kernel.cuh")] * 14 + \
[("PrimitiveOpaqueTriangle.cuh", "ProjectionFwd_kernel.cuh")] * 2 + \
[("PrimitiveVoxel.cuh", "ProjectionFwd_kernel.cuh")] * 2
generate_kernel_instantiation("ProjectionFwd", definition, map_header, map_body, includes)
def generate_ProjectionBwd():
definition = extract_kernel_definition("ProjectionBwd.cu", "projection_fused_bwd_kernel_wrapper")
map_header = ["typename SplatPrimitive", None, None]
map_body = [
("Vanilla3DGS", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("Vanilla3DGS", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("Vanilla3DGS", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::Position"),
("Vanilla3DGS", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::Position"),
("Vanilla3DGS", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::AllReasonable"),
("Vanilla3DGS", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::AllReasonable"),
("MipSplatting", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("MipSplatting", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("MipSplatting", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::Position"),
("MipSplatting", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::Position"),
("MipSplatting", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::AllReasonable"),
("MipSplatting", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::AllReasonable"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::Position"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::Position"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::AllReasonable"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::AllReasonable"),
("SphericalVoronoi3DGUT<2>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<2>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<3>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<3>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<4>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<4>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<5>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<5>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<6>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<6>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<7>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<7>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<8>", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("SphericalVoronoi3DGUT<8>", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("OpaqueTriangle", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("OpaqueTriangle", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
("VoxelPrimitive", "gsplat::CameraModelType::PINHOLE", "HessianDiagonalOutputMode::None"),
("VoxelPrimitive", "gsplat::CameraModelType::FISHEYE", "HessianDiagonalOutputMode::None"),
]
includes = [("Primitive3DGS.cuh", "ProjectionBwd_kernel.cuh")] * 12 + \
[("Primitive3DGUT.cuh", "ProjectionBwd_kernel.cuh")] * 6 + \
[("Primitive3DGUT_SV.cuh", "ProjectionBwd_kernel.cuh")] * 14 + \
[("PrimitiveOpaqueTriangle.cuh", "ProjectionBwd_kernel.cuh")] * 2 + \
[("PrimitiveVoxel.cuh", "ProjectionBwd_kernel.cuh")] * 2
generate_kernel_instantiation("ProjectionBwd", definition, map_header, map_body, includes)
def generate_RasterizationEval3DFwd():
definition = extract_kernel_definition("RasterizationEval3DFwd.cu", "rasterize_to_pixels_eval3d_fwd_kernel_wrapper")
map_header = ["typename SplatPrimitive", None, None, None]
map_body = [
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "false", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "true", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "false", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "true", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::PINHOLE", "false", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::PINHOLE", "true", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::FISHEYE", "false", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::FISHEYE", "true", "false"),
("VoxelPrimitive", "gsplat::CameraModelType::PINHOLE", "true", "true"),
("VoxelPrimitive", "gsplat::CameraModelType::FISHEYE", "true", "true"),
]
includes = [("Primitive3DGUT.cuh", "RasterizationEval3DFwd_kernel.cuh")] * 4 + \
[("Primitive3DGUT_SV.cuh", "RasterizationEval3DFwd_kernel.cuh")] * 4 + \
[("PrimitiveVoxel.cuh", "RasterizationEval3DFwd_kernel.cuh")] * 2
generate_kernel_instantiation("RasterizationEval3DFwd", definition, map_header, map_body, includes)
def generate_RasterizationEval3DBwd():
definition = extract_kernel_definition("RasterizationEval3DBwd.cu", "rasterize_to_pixels_eval3d_bwd_kernel_wrapper")
map_header = ["typename SplatPrimitive", None, None, None, None]
map_body = [
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "false", "false", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "false", "true", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "true", "false", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "true", "true", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "false", "false", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "false", "true", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "true", "false", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "true", "true", "false"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "false", "false", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "false", "true", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "true", "false", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::PINHOLE", "true", "true", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "false", "false", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "false", "true", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "true", "false", "true"),
("Vanilla3DGUT", "gsplat::CameraModelType::FISHEYE", "true", "true", "true"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::PINHOLE", "false", "false", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::PINHOLE", "false", "true", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::PINHOLE", "true", "false", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::PINHOLE", "true", "true", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::FISHEYE", "false", "false", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::FISHEYE", "false", "true", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::FISHEYE", "true", "false", "false"),
("SphericalVoronoi3DGUT_Default", "gsplat::CameraModelType::FISHEYE", "true", "true", "false"),
("VoxelPrimitive", "gsplat::CameraModelType::PINHOLE", "false", "false", "false"),
("VoxelPrimitive", "gsplat::CameraModelType::PINHOLE", "false", "true", "false"),
("VoxelPrimitive", "gsplat::CameraModelType::FISHEYE", "false", "false", "false"),
("VoxelPrimitive", "gsplat::CameraModelType::FISHEYE", "false", "true", "false"),
]
includes = [("Primitive3DGUT.cuh", "RasterizationEval3DBwd_kernel.cuh")] * 16 + \
[("Primitive3DGUT_SV.cuh", "RasterizationEval3DBwd_kernel.cuh")] * 8 + \
[("PrimitiveVoxel.cuh", "RasterizationEval3DBwd_kernel.cuh")] * 4
generate_kernel_instantiation("RasterizationEval3DBwd", definition, map_header, map_body, includes)
generate_ProjectionFwd()
generate_ProjectionBwd()
generate_RasterizationEval3DFwd()
generate_RasterizationEval3DBwd()
+1 -1
View File
@@ -1040,7 +1040,7 @@ class SpirulaeModel(Model):
compute_hessian_diagonal=self.config.compute_hessian_diagonal,
**kwargs,
)
torch.cuda.empty_cache()
# torch.cuda.empty_cache()
if self.config.compute_hessian_diagonal is not None:
rgbd = list(rgbd)
rgbd[0], backward_metadata_injector = _BackwardMetadataInjector.apply(rgbd[0])
+4 -4
View File
@@ -106,18 +106,18 @@ class SpirulaePipeline(VanillaPipeline):
inputs = [((train_inputs[0], val_inputs[0]), (train_inputs[1], val_inputs[1]))]
for i, (camera, batch) in enumerate(inputs):
torch.cuda.empty_cache()
# torch.cuda.empty_cache()
model_outputs = self._model(camera)
torch.cuda.empty_cache()
# torch.cuda.empty_cache()
metrics_dict = self.model.get_metrics_dict(model_outputs, batch)
is_not_last = (i != len(inputs) - 1)
loss_dict = self.model.get_loss_dict(model_outputs, batch, metrics_dict, no_static_losses=is_not_last)
torch.cuda.empty_cache()
# torch.cuda.empty_cache()
if is_not_last:
torch.stack([
x for x in loss_dict.values() if isinstance(x, torch.Tensor)
]).sum().backward()
torch.cuda.empty_cache()
# torch.cuda.empty_cache()
return model_outputs, loss_dict, metrics_dict
@@ -113,10 +113,13 @@ struct _Base3DGS<antialiased>::World {
float opacity;
FixedArray<float3, 16> sh_coeffs;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, std::optional<at::Tensor>> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor means;
at::Tensor quats;
@@ -165,6 +168,7 @@ struct _Base3DGS<antialiased>::World {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ means;
@@ -177,6 +181,7 @@ struct _Base3DGS<antialiased>::World {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.means);
CHECK_INPUT(tensors.means);
@@ -196,6 +201,7 @@ struct _Base3DGS<antialiased>::World {
num_sh = tensors.features_sh.has_value() ?
tensors.features_sh.value().size(-2) : 0;
}
#endif
};
#ifdef __CUDACC__
@@ -256,10 +262,13 @@ struct _Base3DGS<antialiased>::RenderOutput {
float3 rgb;
float depth;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor rgbs;
at::Tensor depths;
@@ -299,6 +308,7 @@ struct _Base3DGS<antialiased>::RenderOutput {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ rgbs;
@@ -306,6 +316,7 @@ struct _Base3DGS<antialiased>::RenderOutput {
Buffer() : rgbs(nullptr), depths(nullptr) {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.rgbs);
CHECK_INPUT(tensors.rgbs);
@@ -313,6 +324,7 @@ struct _Base3DGS<antialiased>::RenderOutput {
rgbs = (float3*)tensors.rgbs.template data_ptr<float>();
depths = tensors.depths.template data_ptr<float>();
}
#endif
};
#ifdef __CUDACC__
@@ -372,11 +384,14 @@ struct _Base3DGS<antialiased>::Screen {
float3 rgb;
float2 xy_abs;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTupleProj;
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, std::optional<at::Tensor>> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor means2d;
at::Tensor depths;
@@ -465,6 +480,7 @@ struct _Base3DGS<antialiased>::Screen {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float2* __restrict__ means2d; // [I, N, 2] or [nnz, 2]
@@ -476,6 +492,7 @@ struct _Base3DGS<antialiased>::Screen {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.means2d);
CHECK_INPUT(tensors.means2d);
@@ -492,6 +509,7 @@ struct _Base3DGS<antialiased>::Screen {
(float2*)tensors.absgrad.value().template data_ptr<float>()
: nullptr;
}
#endif
};
#ifdef __CUDACC__
@@ -85,10 +85,13 @@ struct Vanilla3DGUT::World : public Base3DGUT::World {
FixedArray<float3, 16> sh_coeffs;
#endif
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, std::optional<at::Tensor>> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor means;
at::Tensor quats;
@@ -137,6 +140,7 @@ struct Vanilla3DGUT::World : public Base3DGUT::World {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ means;
@@ -149,6 +153,7 @@ struct Vanilla3DGUT::World : public Base3DGUT::World {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.means);
CHECK_INPUT(tensors.means);
@@ -168,6 +173,7 @@ struct Vanilla3DGUT::World : public Base3DGUT::World {
num_sh = tensors.features_sh.has_value() ?
tensors.features_sh.value().size(-2) : 0;
}
#endif
};
#ifdef __CUDACC__
@@ -88,10 +88,13 @@ struct SphericalVoronoi3DGUT<num_sv>::World : public Base3DGUT::World {
FixedArray<float3, num_sv> sv_colors;
#endif
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor means;
at::Tensor quats;
@@ -138,6 +141,7 @@ struct SphericalVoronoi3DGUT<num_sv>::World : public Base3DGUT::World {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ means;
@@ -149,6 +153,7 @@ struct SphericalVoronoi3DGUT<num_sv>::World : public Base3DGUT::World {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.means);
CHECK_INPUT(tensors.means);
@@ -164,6 +169,7 @@ struct SphericalVoronoi3DGUT<num_sv>::World : public Base3DGUT::World {
sv_sites = (float3*)tensors.sv_sites.template data_ptr<float>();
sv_colors = (float3*)tensors.sv_colors.template data_ptr<float>();
}
#endif
};
#ifdef __CUDACC__
@@ -55,10 +55,13 @@ struct Base3DGUT::RenderOutput {
float3 rgb;
float depth;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor rgbs;
at::Tensor depths;
@@ -98,6 +101,7 @@ struct Base3DGUT::RenderOutput {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ rgbs;
@@ -105,6 +109,7 @@ struct Base3DGUT::RenderOutput {
Buffer() : rgbs(nullptr), depths(nullptr) {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.rgbs);
CHECK_INPUT(tensors.rgbs);
@@ -112,6 +117,7 @@ struct Base3DGUT::RenderOutput {
rgbs = (float3*)tensors.rgbs.data_ptr<float>();
depths = tensors.depths.data_ptr<float>();
}
#endif
};
#ifdef __CUDACC__
@@ -176,11 +182,14 @@ struct Base3DGUT::Screen {
float3x3 iscl_rot;
#endif
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTupleProj;
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
bool hasWorld;
at::Tensor means;
@@ -269,6 +278,7 @@ struct Base3DGUT::Screen {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ means;
@@ -281,6 +291,7 @@ struct Base3DGUT::Screen {
Buffer() {} // uninitialized
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.depths);
if (tensors.hasWorld) {
@@ -300,6 +311,7 @@ struct Base3DGUT::Screen {
size = tensors.hasWorld ?
tensors.quats.numel() / 4 : tensors.opacities.numel();
}
#endif
};
#ifdef __CUDACC__
@@ -103,11 +103,14 @@ struct OpaqueTriangle::World {
FixedArray<float3, 16> sh_coeffs;
FixedArray<float3, 2> ch_coeffs;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTuple;
// typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor means;
at::Tensor quats;
@@ -165,6 +168,7 @@ struct OpaqueTriangle::World {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ means;
@@ -179,6 +183,7 @@ struct OpaqueTriangle::World {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.means);
CHECK_INPUT(tensors.means);
@@ -200,6 +205,7 @@ struct OpaqueTriangle::World {
num_sh = tensors.features_sh.size(-2);
features_ch = (float3*)tensors.features_ch.data_ptr<float>();
}
#endif
};
#ifdef __CUDACC__
@@ -277,10 +283,13 @@ struct OpaqueTriangle::RenderOutput {
float depth;
float3 normal;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor rgbs;
at::Tensor depths;
@@ -323,6 +332,7 @@ struct OpaqueTriangle::RenderOutput {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ rgbs;
@@ -331,6 +341,7 @@ struct OpaqueTriangle::RenderOutput {
Buffer() : rgbs(nullptr), depths(nullptr), normals(nullptr) {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.rgbs);
CHECK_INPUT(tensors.rgbs);
@@ -340,6 +351,7 @@ struct OpaqueTriangle::RenderOutput {
depths = tensors.depths.data_ptr<float>();
normals = (float3*)tensors.normals.data_ptr<float>();
}
#endif
};
#ifdef __CUDACC__
@@ -407,11 +419,14 @@ struct OpaqueTriangle::Screen {
#endif
float3 normal;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTupleProj;
typedef std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
bool hasWorld;
at::Tensor hardness;
@@ -496,6 +511,7 @@ struct OpaqueTriangle::Screen {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float2* __restrict__ hardness = nullptr;
@@ -507,6 +523,7 @@ struct OpaqueTriangle::Screen {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.verts);
if (tensors.hasWorld) {
@@ -523,6 +540,7 @@ struct OpaqueTriangle::Screen {
size = tensors.hasWorld ?
tensors.hardness.numel() / 2 : tensors.verts.numel() / 9;
}
#endif
};
#ifdef __CUDACC__
@@ -100,10 +100,13 @@ struct VoxelPrimitive::World {
FixedArray<float, 8> densities;
FixedArray<float3, 16> sh_coeffs;
#ifndef NO_TORCH
typedef std::tuple<std::optional<at::Tensor>, std::optional<at::Tensor>, at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
std::optional<at::Tensor> pos_size;
std::optional<at::Tensor> densities;
@@ -144,6 +147,7 @@ struct VoxelPrimitive::World {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float4* __restrict__ pos_size;
@@ -154,6 +158,7 @@ struct VoxelPrimitive::World {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.features_dc);
if (tensors.pos_size.has_value())
@@ -170,6 +175,7 @@ struct VoxelPrimitive::World {
features_sh = (float3*)tensors.features_sh.data_ptr<float>();
num_sh = tensors.features_sh.size(-2);
}
#endif
};
#ifdef __CUDACC__
@@ -230,10 +236,13 @@ struct VoxelPrimitive::RenderOutput {
float3 rgb;
float depth;
#ifndef NO_TORCH
typedef std::tuple<at::Tensor, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
at::Tensor rgbs;
at::Tensor depths;
@@ -273,6 +282,7 @@ struct VoxelPrimitive::RenderOutput {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float3* __restrict__ rgbs;
@@ -280,6 +290,7 @@ struct VoxelPrimitive::RenderOutput {
Buffer() : rgbs(nullptr), depths(nullptr) {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.rgbs);
CHECK_INPUT(tensors.rgbs);
@@ -287,6 +298,7 @@ struct VoxelPrimitive::RenderOutput {
rgbs = (float3*)tensors.rgbs.data_ptr<float>();
depths = tensors.depths.data_ptr<float>();
}
#endif
};
#ifdef __CUDACC__
@@ -346,11 +358,14 @@ struct VoxelPrimitive::Screen {
float3 rgb;
float density_abs;
#ifndef NO_TORCH
typedef std::tuple<std::optional<at::Tensor>, at::Tensor> TensorTupleProj;
typedef std::tuple<std::optional<at::Tensor>, std::optional<at::Tensor>, std::optional<at::Tensor>, at::Tensor> TensorTuple;
#endif
struct Buffer;
#ifndef NO_TORCH
struct Tensor {
bool hasWorld;
std::optional<at::Tensor> pos_size;
@@ -428,6 +443,7 @@ struct VoxelPrimitive::Screen {
Buffer buffer() { return Buffer(*this); }
};
#endif
struct Buffer {
float4* __restrict__ pos_size = nullptr;
@@ -438,6 +454,7 @@ struct VoxelPrimitive::Screen {
Buffer() {}
#ifndef NO_TORCH
Buffer(const Tensor& tensors) {
DEVICE_GUARD(tensors.rgbs);
if (tensors.pos_size.has_value()) {
@@ -456,6 +473,7 @@ struct VoxelPrimitive::Screen {
rgbs = (float3*)tensors.rgbs.data_ptr<float>();
size = tensors.rgbs.size(-2);
}
#endif
};
#ifdef __CUDACC__
@@ -6,51 +6,6 @@
#include "Primitive3DGS.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
Vanilla3DGS::Screen::TensorTupleProj // out splats
> projection_3dgs_forward_tensor(
// inputs
const Vanilla3DGS::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<Vanilla3DGS>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgs_backward_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<Vanilla3DGS>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
@@ -1,31 +0,0 @@
#include "ProjectionBwd.cuh"
#include "Primitive3DGS.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGS, HessianDiagonalOutputMode::AllReasonable>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
@@ -1,31 +0,0 @@
#include "ProjectionBwd.cuh"
#include "Primitive3DGS.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGS, HessianDiagonalOutputMode::Position>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
@@ -6,51 +6,6 @@
#include "Primitive3DGUT.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
Vanilla3DGUT::Screen::TensorTupleProj // out splats
> projection_3dgut_forward_tensor(
// inputs
const Vanilla3DGUT::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<Vanilla3DGUT>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_backward_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<Vanilla3DGUT>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
@@ -1,32 +0,0 @@
#include "ProjectionBwd.cuh"
#include "Primitive3DGUT.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGUT, HessianDiagonalOutputMode::AllReasonable>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
@@ -1,32 +0,0 @@
#include "ProjectionBwd.cuh"
#include "Primitive3DGUT.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGUT, HessianDiagonalOutputMode::Position>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
@@ -1,60 +0,0 @@
#include "ProjectionFwd.cuh"
#include "ProjectionBwd.cuh"
#include "Primitive3DGUT_SV.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj // out splats
> projection_3dgut_sv_forward_tensor(
// inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
int num_sv = std::get<5>(in_splats).size(-2);
#define _CASE(n) \
if (num_sv == n) return launch_projection_fused_fwd_kernel<SphericalVoronoi3DGUT<n>>( \
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs); \
_CASE(2) _CASE(3) _CASE(4) _CASE(5) _CASE(6) _CASE(7) _CASE(8)
#undef _CASE
throw std::invalid_argument("Unsupported num_sv");
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
SphericalVoronoi3DGUT_Default::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_sv_backward_tensor(
// fwd inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
int num_sv = std::get<5>(splats_world).size(-2);
#define _CASE(n) \
if (num_sv == n) return launch_projection_projection_fused_bwd_kernel<SphericalVoronoi3DGUT<n>>( \
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs, \
aabb, v_splats_screen, viewmats_requires_grad);
_CASE(2) _CASE(3) _CASE(4) _CASE(5) _CASE(6) _CASE(7) _CASE(8)
#undef _CASE
throw std::invalid_argument("Unsupported num_sv");
}
+22 -388
View File
@@ -19,63 +19,50 @@
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // camera_ids
at::Tensor, // gaussian_ids
at::Tensor, // aabb
Vanilla3DGS::Screen::TensorTupleProj // out splats
> projection_3dgs_forward_tensor(
Vanilla3DGUT::Screen::TensorTuple // out splats
> projection_3dgut_hetero_forward_tensor(
// inputs
const Vanilla3DGS::World::TensorTuple &in_splats,
const Vanilla3DGUT::World::TensorTuple &in_splats_tensor,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
const CameraDistortionCoeffsTensor dist_coeffs,
const at::Tensor intersection_count_map, // [C+1]
const at::Tensor intersection_splat_id // [nnz]
);
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgs_backward_tensor(
> projection_3dgut_hetero_backward_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const Vanilla3DGUT::World::TensorTuple &splats_world_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
const at::Tensor camera_ids, // [nnz]
const at::Tensor gaussian_ids, // [nnz]
const at::Tensor aabb, // [nnz, 4]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
const Vanilla3DGUT::Screen::TensorTuple &v_splats_proj_tuple,
const bool viewmats_requires_grad,
const bool sparse_grad
);
@@ -127,237 +114,6 @@ std::tuple<
);
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // aabb
Vanilla3DGUT::Screen::TensorTupleProj // out splats
> projection_3dgut_forward_tensor(
// inputs
const Vanilla3DGUT::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_backward_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // camera_ids
at::Tensor, // gaussian_ids
at::Tensor, // aabb
Vanilla3DGUT::Screen::TensorTuple // out splats
> projection_3dgut_hetero_forward_tensor(
// inputs
const Vanilla3DGUT::World::TensorTuple &in_splats_tensor,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const at::Tensor intersection_count_map, // [C+1]
const at::Tensor intersection_splat_id // [nnz]
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_hetero_backward_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor camera_ids, // [nnz]
const at::Tensor gaussian_ids, // [nnz]
const at::Tensor aabb, // [nnz, 4]
// grad outputs
const Vanilla3DGUT::Screen::TensorTuple &v_splats_proj_tuple,
const bool viewmats_requires_grad,
const bool sparse_grad
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // aabb
SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj // out splats
> projection_3dgut_sv_forward_tensor(
// inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
SphericalVoronoi3DGUT_Default::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_sv_backward_tensor(
// fwd inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // aabb
MipSplatting::Screen::TensorTupleProj // out splats
> projection_mip_forward_tensor(
// inputs
const MipSplatting::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_mip_backward_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // camera_ids
at::Tensor, // gaussian_ids
@@ -406,91 +162,6 @@ std::tuple<
);
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // aabb
OpaqueTriangle::Screen::TensorTupleProj // out splats
> projection_opaque_triangle_forward_tensor(
// inputs
const OpaqueTriangle::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
OpaqueTriangle::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_opaque_triangle_backward_tensor(
// fwd inputs
const OpaqueTriangle::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const OpaqueTriangle::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
at::Tensor, // camera_ids
at::Tensor, // gaussian_ids
@@ -537,40 +208,3 @@ std::tuple<
const bool viewmats_requires_grad,
const bool sparse_grad
);
std::tuple<
at::Tensor, // aabb
VoxelPrimitive::Screen::TensorTupleProj // out splats
> projection_voxel_forward_tensor(
// inputs
const VoxelPrimitive::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
VoxelPrimitive::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_voxel_backward_tensor(
// fwd inputs
const VoxelPrimitive::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const VoxelPrimitive::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
@@ -5,50 +5,6 @@
#include "Primitive3DGS.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
MipSplatting::Screen::TensorTupleProj // out splats
> projection_mip_forward_tensor(
// inputs
const MipSplatting::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<MipSplatting>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_mip_backward_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<MipSplatting>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
@@ -1,31 +0,0 @@
#include "ProjectionBwd.cuh"
#include "Primitive3DGS.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<MipSplatting, HessianDiagonalOutputMode::AllReasonable>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
@@ -1,31 +0,0 @@
#include "ProjectionBwd.cuh"
#include "Primitive3DGS.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<MipSplatting, HessianDiagonalOutputMode::Position>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
@@ -5,50 +5,6 @@
#include "PrimitiveOpaqueTriangle.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
OpaqueTriangle::Screen::TensorTupleProj // out splats
> projection_opaque_triangle_forward_tensor(
// inputs
const OpaqueTriangle::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<OpaqueTriangle>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
OpaqueTriangle::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_opaque_triangle_backward_tensor(
// fwd inputs
const OpaqueTriangle::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const OpaqueTriangle::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<OpaqueTriangle>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // camera_ids
@@ -1,50 +0,0 @@
#include "ProjectionFwd.cuh"
#include "ProjectionBwd.cuh"
#include "PrimitiveVoxel.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
VoxelPrimitive::Screen::TensorTupleProj // out splats
> projection_voxel_forward_tensor(
// inputs
const VoxelPrimitive::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<VoxelPrimitive>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
VoxelPrimitive::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_voxel_backward_tensor(
// fwd inputs
const VoxelPrimitive::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const VoxelPrimitive::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<VoxelPrimitive>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
@@ -0,0 +1,524 @@
#include "ProjectionBwd.cuh"
#include <gsplat/Utils.cuh>
#include <c10/cuda/CUDAStream.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
template<
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
HessianDiagonalOutputMode hessian_diagonal_output_mode
>
void projection_fused_bwd_kernel_wrapper(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
typename SplatPrimitive::Screen::Buffer v_splats_screen,
typename SplatPrimitive::Screen::Buffer vr_splats_screen,
typename SplatPrimitive::Screen::Buffer h_splats_screen,
// grad inputs
typename SplatPrimitive::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
typename SplatPrimitive::World::Buffer vr_splats_world,
typename SplatPrimitive::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
template<typename SplatPrimitive, HessianDiagonalOutputMode hessian_diagonal_output_mode>
inline std::tuple<
typename SplatPrimitive::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, typename SplatPrimitive::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, typename SplatPrimitive::World::TensorTuple> // h_world_pos or h_splats
> _launch_projection_projection_fused_bwd_kernel(
// fwd inputs
const typename SplatPrimitive::World::TensorTuple &splats_world_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const typename SplatPrimitive::Screen::TensorTupleProj &v_splats_screen_tuple,
const typename SplatPrimitive::Screen::TensorTupleProj *vr_splats_screen_tuple,
const typename SplatPrimitive::Screen::TensorTupleProj *h_splats_screen_tuple,
const bool viewmats_requires_grad
) {
typename SplatPrimitive::World::Tensor splats_world(splats_world_tuple);
uint32_t N = splats_world.size(); // number of gaussians
uint32_t C = viewmats.size(-3); // number of cameras
uint32_t B = splats_world.batchSize(); // number of batches
typename SplatPrimitive::Screen::Tensor v_splats_screen(v_splats_screen_tuple);
typename SplatPrimitive::World::Tensor v_splats_world = splats_world.allocProjBwd(false);
typename SplatPrimitive::Screen::Tensor vr_splats_screen;
typename SplatPrimitive::Screen::Tensor h_splats_screen;
if (hessian_diagonal_output_mode != HessianDiagonalOutputMode::None) {
vr_splats_screen = *vr_splats_screen_tuple;
h_splats_screen = *h_splats_screen_tuple;
}
auto opt = splats_world.options();
at::Tensor v_viewmats;
if (viewmats_requires_grad)
v_viewmats = zeros_like<float>(viewmats);
at::Tensor vr_world_pos;
at::Tensor h_world_pos;
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position) {
vr_world_pos = at::empty({B, N, 3}, opt);
h_world_pos = at::empty({B, N, 3}, opt);
set_zero<float>(vr_world_pos);
set_zero<float>(h_world_pos);
}
typename SplatPrimitive::World::Tensor vr_splats_world;
typename SplatPrimitive::World::Tensor h_splats_world;
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable) {
vr_splats_world = splats_world.allocProjBwd(true);
h_splats_world = splats_world.allocProjBwd(true);
}
#define _LAUNCH_ARGS ( \
(cudaStream_t)at::cuda::getCurrentCUDAStream(), B, C, N, \
splats_world.buffer(), viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
image_width, image_height, (int4*)aabb.data_ptr<int32_t>(), \
v_splats_screen.buffer(), \
hessian_diagonal_output_mode != HessianDiagonalOutputMode::None ? vr_splats_screen.buffer() : typename SplatPrimitive::Screen::Buffer{}, \
hessian_diagonal_output_mode != HessianDiagonalOutputMode::None ? h_splats_screen.buffer() : typename SplatPrimitive::Screen::Buffer{}, \
v_splats_world.buffer(), \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position ? (float3*)vr_world_pos.data_ptr<float>() : nullptr, \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position ? (float3*)h_world_pos.data_ptr<float>() : nullptr, \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable ? vr_splats_world.buffer() : typename SplatPrimitive::World::Buffer{}, \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable ? h_splats_world.buffer() : typename SplatPrimitive::World::Buffer{}, \
viewmats_requires_grad ? v_viewmats.data_ptr<float>() : nullptr \
)
if (camera_model == gsplat::CameraModelType::PINHOLE)
projection_fused_bwd_kernel_wrapper<SplatPrimitive, gsplat::CameraModelType::PINHOLE, hessian_diagonal_output_mode> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
projection_fused_bwd_kernel_wrapper<SplatPrimitive, gsplat::CameraModelType::FISHEYE, hessian_diagonal_output_mode> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable)
return std::make_tuple(v_splats_world.tuple(), v_viewmats, vr_splats_world.tuple(), h_splats_world.tuple());
return std::make_tuple(v_splats_world.tuple(), v_viewmats, vr_world_pos, h_world_pos);
}
template<typename SplatPrimitive>
inline std::tuple<
typename SplatPrimitive::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> launch_projection_projection_fused_bwd_kernel(
// fwd inputs
const typename SplatPrimitive::World::TensorTuple &splats_world_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const typename SplatPrimitive::Screen::TensorTupleProj &v_splats_screen_tuple,
const bool viewmats_requires_grad
) {
auto [v_splats, v_viewmats, vr_splats, h_splats] =
_launch_projection_projection_fused_bwd_kernel
<SplatPrimitive, HessianDiagonalOutputMode::None>
(
splats_world_tuple,
viewmats,
intrins,
image_width,
image_height,
camera_model,
dist_coeffs,
aabb,
v_splats_screen_tuple,
nullptr,
nullptr,
viewmats_requires_grad
);
return std::make_tuple(v_splats, v_viewmats);
}
// ================
// Vanilla3DGS
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgs_backward_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<Vanilla3DGS>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGS, HessianDiagonalOutputMode::AllReasonable>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGS, HessianDiagonalOutputMode::Position>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
// ================
// MipSplatting
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_mip_backward_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<MipSplatting>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<MipSplatting, HessianDiagonalOutputMode::AllReasonable>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<MipSplatting, HessianDiagonalOutputMode::Position>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
// ================
// Vanilla3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_backward_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<Vanilla3DGUT>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGUT, HessianDiagonalOutputMode::AllReasonable>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
return _launch_projection_projection_fused_bwd_kernel<Vanilla3DGUT, HessianDiagonalOutputMode::Position>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, &vr_splats_screen, &h_splats_screen, viewmats_requires_grad);
}
// ================
// SphericalVoronoi3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
SphericalVoronoi3DGUT_Default::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_sv_backward_tensor(
// fwd inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
int num_sv = std::get<5>(splats_world).size(-2);
#define _CASE(n) \
if (num_sv == n) return launch_projection_projection_fused_bwd_kernel<SphericalVoronoi3DGUT<n>>( \
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs, \
aabb, v_splats_screen, viewmats_requires_grad);
_CASE(2) _CASE(3) _CASE(4) _CASE(5) _CASE(6) _CASE(7) _CASE(8)
#undef _CASE
throw std::invalid_argument("Unsupported num_sv");
}
// ================
// OpaqueTriangle
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
OpaqueTriangle::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_opaque_triangle_backward_tensor(
// fwd inputs
const OpaqueTriangle::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const OpaqueTriangle::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<OpaqueTriangle>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
// ================
// VoxelPrimitive
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
VoxelPrimitive::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_voxel_backward_tensor(
// fwd inputs
const VoxelPrimitive::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const VoxelPrimitive::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
) {
return launch_projection_projection_fused_bwd_kernel<VoxelPrimitive>(
splats_world, viewmats, intrins, image_width, image_height, camera_model, dist_coeffs,
aabb, v_splats_screen, viewmats_requires_grad);
}
+257 -275
View File
@@ -6,270 +6,49 @@
#include <ATen/Tensor.h>
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#include "Primitive3DGS.cuh"
#include "Primitive3DGUT.cuh"
#include "Primitive3DGUT_SV.cuh"
#include "PrimitiveOpaqueTriangle.cuh"
#include "PrimitiveVoxel.cuh"
#include "types.cuh"
#include <c10/cuda/CUDAStream.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
#include "common.cuh"
enum class HessianDiagonalOutputMode {
None,
Position,
AllReasonable
};
template<
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
HessianDiagonalOutputMode hessian_diagonal_output_mode
>
__global__ void projection_fused_bwd_kernel(
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 *__restrict__ aabb, // [B, C, N, 4]
// grad outputs
typename SplatPrimitive::Screen::Buffer v_splats_screen,
typename SplatPrimitive::Screen::Buffer vr_splats_screen,
typename SplatPrimitive::Screen::Buffer h_splats_screen,
// grad inputs
typename SplatPrimitive::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
typename SplatPrimitive::World::Buffer vr_splats_world,
typename SplatPrimitive::World::Buffer h_splats_world,
float *__restrict__ v_viewmats // [B, C, 4, 4] optional
) {
// parallelize over B * C * N.
uint32_t idx = cg::this_grid().thread_rank();
if (idx >= B * C * N || (aabb[idx].z-aabb[idx].x)*(aabb[idx].w-aabb[idx].y) <= 0) {
return;
}
const uint32_t bid = idx / (C * N); // batch id
const uint32_t cid = (idx / N) % C; // camera id
const uint32_t gid = idx % N; // gaussian id
// Load camera
viewmats += bid * C * 16 + cid * 16;
float4 intrin = intrins[bid * C + cid];
float3x3 R = {
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
typename SplatPrimitive::BwdProjCamera cam = {
R, t, fx, fy, cx, cy,
image_width, image_height,
};
cam.dist_coeffs = dist_coeffs_buffer.load(bid * C + cid);
// Load splat
typename SplatPrimitive::World splat_world =
SplatPrimitive::World::load(splats_world, bid * N + gid);
typename SplatPrimitive::Screen v_splat_screen =
SplatPrimitive::Screen::load(v_splats_screen, idx);
typename SplatPrimitive::Screen vr_splat_screen;
typename SplatPrimitive::Screen h_splat_screen;
if (hessian_diagonal_output_mode != HessianDiagonalOutputMode::None) {
vr_splat_screen = SplatPrimitive::Screen::load(vr_splats_screen, idx);
h_splat_screen = SplatPrimitive::Screen::load(h_splats_screen, idx);
}
// Projection
typename SplatPrimitive::World v_splat_world = SplatPrimitive::World::zero();
float3x3 v_R = {0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f};
float3 v_t = {0.f, 0.f, 0.f};
float3 vr_world_pos = {0.f, 0.f, 0.f};
float3 h_world_pos = {0.f, 0.f, 0.f};
typename SplatPrimitive::World vr_splat_world = SplatPrimitive::World::zero();
typename SplatPrimitive::World h_splat_world = SplatPrimitive::World::zero();
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::None) {
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp_vjp(splat_world, cam, v_splat_screen, v_splat_world, v_R, v_t);
break;
// case gsplat::CameraModelType::ORTHO: // orthographic projection
// SplatPrimitive::project_ortho_vjp(splat_world, cam, v_splat_screen, v_splat_world, v_R, v_t);
// break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye_vjp(splat_world, cam, v_splat_screen, v_splat_world, v_R, v_t);
break;
}
} else if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position) {
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_world_pos, h_world_pos);
break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_world_pos, h_world_pos);
break;
}
} else if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable) {
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_splat_world, h_splat_world);
break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_splat_world, h_splat_world);
break;
}
}
// Save results
v_splat_world.atomicAddGradientToBuffer(v_splats_world, bid * N + gid);
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position) {
atomicAddFVec(&vr_world_pos_buffer[bid * N + gid], vr_world_pos);
atomicAddFVec(&h_world_pos_buffer[bid * N + gid], h_world_pos);
} else if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable) {
vr_splat_world.atomicAddGradientToBuffer(vr_splats_world, bid * N + gid);
h_splat_world.atomicAddGradientToBuffer(h_splats_world, bid * N + gid);
}
if (v_viewmats != nullptr) {
auto warp = cg::tiled_partition<WARP_SIZE>(cg::this_thread_block());
auto warp_group_c = cg::labeled_partition(warp, cid);
warpSum(v_R[0], warp_group_c);
warpSum(v_R[1], warp_group_c);
warpSum(v_R[2], warp_group_c);
warpSum(v_t, warp_group_c);
if (warp_group_c.thread_rank() == 0) {
v_viewmats += bid * C * 16 + cid * 16;
#pragma unroll
for (uint32_t i = 0; i < 3; i++) { // rows
atomicAdd(v_viewmats + i * 4 + 0, v_R[i].x);
atomicAdd(v_viewmats + i * 4 + 1, v_R[i].y);
atomicAdd(v_viewmats + i * 4 + 2, v_R[i].z);
}
atomicAdd(v_viewmats + 0 * 4 + 3, v_t.x);
atomicAdd(v_viewmats + 1 * 4 + 3, v_t.y);
atomicAdd(v_viewmats + 2 * 4 + 3, v_t.z);
}
}
}
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
template<typename SplatPrimitive, HessianDiagonalOutputMode hessian_diagonal_output_mode>
inline std::tuple<
typename SplatPrimitive::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, typename SplatPrimitive::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, typename SplatPrimitive::World::TensorTuple> // h_world_pos or h_splats
> _launch_projection_projection_fused_bwd_kernel(
// fwd inputs
const typename SplatPrimitive::World::TensorTuple &splats_world_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const typename SplatPrimitive::Screen::TensorTupleProj &v_splats_screen_tuple,
const typename SplatPrimitive::Screen::TensorTupleProj *vr_splats_screen_tuple,
const typename SplatPrimitive::Screen::TensorTupleProj *h_splats_screen_tuple,
const bool viewmats_requires_grad
) {
typename SplatPrimitive::World::Tensor splats_world(splats_world_tuple);
uint32_t N = splats_world.size(); // number of gaussians
uint32_t C = viewmats.size(-3); // number of cameras
uint32_t B = splats_world.batchSize(); // number of batches
typename SplatPrimitive::Screen::Tensor v_splats_screen(v_splats_screen_tuple);
typename SplatPrimitive::World::Tensor v_splats_world = splats_world.allocProjBwd(false);
typename SplatPrimitive::Screen::Tensor vr_splats_screen;
typename SplatPrimitive::Screen::Tensor h_splats_screen;
if (hessian_diagonal_output_mode != HessianDiagonalOutputMode::None) {
vr_splats_screen = *vr_splats_screen_tuple;
h_splats_screen = *h_splats_screen_tuple;
}
auto opt = splats_world.options();
at::Tensor v_viewmats;
if (viewmats_requires_grad)
v_viewmats = zeros_like<float>(viewmats);
at::Tensor vr_world_pos;
at::Tensor h_world_pos;
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position) {
vr_world_pos = at::empty({B, N, 3}, opt);
h_world_pos = at::empty({B, N, 3}, opt);
set_zero<float>(vr_world_pos);
set_zero<float>(h_world_pos);
}
typename SplatPrimitive::World::Tensor vr_splats_world;
typename SplatPrimitive::World::Tensor h_splats_world;
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable) {
vr_splats_world = splats_world.allocProjBwd(true);
h_splats_world = splats_world.allocProjBwd(true);
}
#define _LAUNCH_ARGS \
<<<_LAUNCH_ARGS_1D(B*C*N, block)>>>( \
B, C, N, \
splats_world.buffer(), viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
image_width, image_height, (int4*)aabb.data_ptr<int32_t>(), \
v_splats_screen.buffer(), \
hessian_diagonal_output_mode != HessianDiagonalOutputMode::None ? vr_splats_screen.buffer() : typename SplatPrimitive::Screen::Buffer{}, \
hessian_diagonal_output_mode != HessianDiagonalOutputMode::None ? h_splats_screen.buffer() : typename SplatPrimitive::Screen::Buffer{}, \
v_splats_world.buffer(), \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position ? (float3*)vr_world_pos.data_ptr<float>() : nullptr, \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position ? (float3*)h_world_pos.data_ptr<float>() : nullptr, \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable ? vr_splats_world.buffer() : typename SplatPrimitive::World::Buffer{}, \
hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable ? h_splats_world.buffer() : typename SplatPrimitive::World::Buffer{}, \
viewmats_requires_grad ? v_viewmats.data_ptr<float>() : nullptr \
)
constexpr uint block = hessian_diagonal_output_mode == HessianDiagonalOutputMode::None ? 128 : 64;
if (camera_model == gsplat::CameraModelType::PINHOLE)
projection_fused_bwd_kernel<SplatPrimitive, gsplat::CameraModelType::PINHOLE, hessian_diagonal_output_mode> _LAUNCH_ARGS;
// else if (camera_model == gsplat::CameraModelType::ORTHO)
// projection_fused_bwd_kernel<SplatPrimitive, gsplat::CameraModelType::ORTHO> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
projection_fused_bwd_kernel<SplatPrimitive, gsplat::CameraModelType::FISHEYE, hessian_diagonal_output_mode> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable)
return std::make_tuple(v_splats_world.tuple(), v_viewmats, vr_splats_world.tuple(), h_splats_world.tuple());
return std::make_tuple(v_splats_world.tuple(), v_viewmats, vr_world_pos, h_world_pos);
}
template<typename SplatPrimitive>
inline std::tuple<
typename SplatPrimitive::World::TensorTuple, // v_splats
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> launch_projection_projection_fused_bwd_kernel(
> projection_3dgs_backward_tensor(
// fwd inputs
const typename SplatPrimitive::World::TensorTuple &splats_world_tuple,
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
@@ -279,25 +58,228 @@ inline std::tuple<
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const typename SplatPrimitive::Screen::TensorTupleProj &v_splats_screen_tuple,
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
) {
auto [v_splats, v_viewmats, vr_splats, h_splats] =
_launch_projection_projection_fused_bwd_kernel
<SplatPrimitive, HessianDiagonalOutputMode::None>
(
splats_world_tuple,
viewmats,
intrins,
image_width,
image_height,
camera_model,
dist_coeffs,
aabb,
v_splats_screen_tuple,
nullptr,
nullptr,
viewmats_requires_grad
);
return std::make_tuple(v_splats, v_viewmats);
}
);
std::tuple<
Vanilla3DGS::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGS::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgs_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGS::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGS::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGS::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_mip_backward_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
MipSplatting::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, MipSplatting::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, MipSplatting::World::TensorTuple> // h_world_pos or h_splats
> projection_mip_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const MipSplatting::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const MipSplatting::Screen::TensorTupleProj &v_splats_screen,
const MipSplatting::Screen::TensorTupleProj &vr_splats_screen,
const MipSplatting::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_backward_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
Vanilla3DGUT::World::TensorTuple, // v_splats
at::Tensor, // v_viewmats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple>, // vr_world_pos or vr_splats
std::variant<at::Tensor, Vanilla3DGUT::World::TensorTuple> // h_world_pos or h_splats
> projection_3dgut_backward_with_position_hessian_diagonal_tensor(
// fwd inputs
const Vanilla3DGUT::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const Vanilla3DGUT::Screen::TensorTupleProj &v_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &vr_splats_screen,
const Vanilla3DGUT::Screen::TensorTupleProj &h_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
SphericalVoronoi3DGUT_Default::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_3dgut_sv_backward_tensor(
// fwd inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
OpaqueTriangle::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_opaque_triangle_backward_tensor(
// fwd inputs
const OpaqueTriangle::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const OpaqueTriangle::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
std::tuple<
VoxelPrimitive::World::TensorTuple, // v_splats
at::Tensor // v_viewmats
> projection_voxel_backward_tensor(
// fwd inputs
const VoxelPrimitive::World::TensorTuple &splats_world,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
// fwd outputs
const at::Tensor aabb, // [..., C, N, 2]
// grad outputs
const VoxelPrimitive::Screen::TensorTupleProj &v_splats_screen,
const bool viewmats_requires_grad
);
@@ -0,0 +1,201 @@
#include <cuda_runtime.h>
#include <cstdint>
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#ifndef NO_TORCH
#define NO_TORCH
#endif
#include "types.cuh"
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
template<
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
HessianDiagonalOutputMode hessian_diagonal_output_mode
>
__global__ void projection_fused_bwd_kernel(
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 *__restrict__ aabb, // [B, C, N, 4]
// grad outputs
typename SplatPrimitive::Screen::Buffer v_splats_screen,
typename SplatPrimitive::Screen::Buffer vr_splats_screen,
typename SplatPrimitive::Screen::Buffer h_splats_screen,
// grad inputs
typename SplatPrimitive::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
typename SplatPrimitive::World::Buffer vr_splats_world,
typename SplatPrimitive::World::Buffer h_splats_world,
float *__restrict__ v_viewmats // [B, C, 4, 4] optional
) {
// parallelize over B * C * N.
uint32_t idx = cg::this_grid().thread_rank();
if (idx >= B * C * N || (aabb[idx].z-aabb[idx].x)*(aabb[idx].w-aabb[idx].y) <= 0) {
return;
}
const uint32_t bid = idx / (C * N); // batch id
const uint32_t cid = (idx / N) % C; // camera id
const uint32_t gid = idx % N; // gaussian id
// Load camera
viewmats += bid * C * 16 + cid * 16;
float4 intrin = intrins[bid * C + cid];
float3x3 R = {
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
typename SplatPrimitive::BwdProjCamera cam = {
R, t, fx, fy, cx, cy,
image_width, image_height,
};
cam.dist_coeffs = dist_coeffs_buffer.load(bid * C + cid);
// Load splat
typename SplatPrimitive::World splat_world =
SplatPrimitive::World::load(splats_world, bid * N + gid);
typename SplatPrimitive::Screen v_splat_screen =
SplatPrimitive::Screen::load(v_splats_screen, idx);
typename SplatPrimitive::Screen vr_splat_screen;
typename SplatPrimitive::Screen h_splat_screen;
if (hessian_diagonal_output_mode != HessianDiagonalOutputMode::None) {
vr_splat_screen = SplatPrimitive::Screen::load(vr_splats_screen, idx);
h_splat_screen = SplatPrimitive::Screen::load(h_splats_screen, idx);
}
// Projection
typename SplatPrimitive::World v_splat_world = SplatPrimitive::World::zero();
float3x3 v_R = {0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f, 0.f};
float3 v_t = {0.f, 0.f, 0.f};
float3 vr_world_pos = {0.f, 0.f, 0.f};
float3 h_world_pos = {0.f, 0.f, 0.f};
typename SplatPrimitive::World vr_splat_world = SplatPrimitive::World::zero();
typename SplatPrimitive::World h_splat_world = SplatPrimitive::World::zero();
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::None) {
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp_vjp(splat_world, cam, v_splat_screen, v_splat_world, v_R, v_t);
break;
// case gsplat::CameraModelType::ORTHO: // orthographic projection
// SplatPrimitive::project_ortho_vjp(splat_world, cam, v_splat_screen, v_splat_world, v_R, v_t);
// break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye_vjp(splat_world, cam, v_splat_screen, v_splat_world, v_R, v_t);
break;
}
} else if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position) {
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_world_pos, h_world_pos);
break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_world_pos, h_world_pos);
break;
}
} else if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable) {
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_splat_world, h_splat_world);
break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye_vjp(splat_world, cam, v_splat_screen,
vr_splat_screen, h_splat_screen, v_splat_world, v_R, v_t, vr_splat_world, h_splat_world);
break;
}
}
// Save results
v_splat_world.atomicAddGradientToBuffer(v_splats_world, bid * N + gid);
if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::Position) {
atomicAddFVec(&vr_world_pos_buffer[bid * N + gid], vr_world_pos);
atomicAddFVec(&h_world_pos_buffer[bid * N + gid], h_world_pos);
} else if (hessian_diagonal_output_mode == HessianDiagonalOutputMode::AllReasonable) {
vr_splat_world.atomicAddGradientToBuffer(vr_splats_world, bid * N + gid);
h_splat_world.atomicAddGradientToBuffer(h_splats_world, bid * N + gid);
}
if (v_viewmats != nullptr) {
auto warp = cg::tiled_partition<WARP_SIZE>(cg::this_thread_block());
auto warp_group_c = cg::labeled_partition(warp, cid);
warpSum(v_R[0], warp_group_c);
warpSum(v_R[1], warp_group_c);
warpSum(v_R[2], warp_group_c);
warpSum(v_t, warp_group_c);
if (warp_group_c.thread_rank() == 0) {
v_viewmats += bid * C * 16 + cid * 16;
#pragma unroll
for (uint32_t i = 0; i < 3; i++) { // rows
atomicAdd(v_viewmats + i * 4 + 0, v_R[i].x);
atomicAdd(v_viewmats + i * 4 + 1, v_R[i].y);
atomicAdd(v_viewmats + i * 4 + 2, v_R[i].z);
}
atomicAdd(v_viewmats + 0 * 4 + 3, v_t.x);
atomicAdd(v_viewmats + 1 * 4 + 3, v_t.y);
atomicAdd(v_viewmats + 2 * 4 + 3, v_t.z);
}
}
}
template<
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
HessianDiagonalOutputMode hessian_diagonal_output_mode
>
void projection_fused_bwd_kernel_wrapper(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
typename SplatPrimitive::Screen::Buffer v_splats_screen,
typename SplatPrimitive::Screen::Buffer vr_splats_screen,
typename SplatPrimitive::Screen::Buffer h_splats_screen,
// grad inputs
typename SplatPrimitive::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
typename SplatPrimitive::World::Buffer vr_splats_world,
typename SplatPrimitive::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
) {
constexpr uint block = hessian_diagonal_output_mode == HessianDiagonalOutputMode::None ? 128 : WARP_SIZE;
projection_fused_bwd_kernel<SplatPrimitive, camera_model, hessian_diagonal_output_mode>
<<<_CEIL_DIV(B*C*N, block), block, 0, stream>>>(
B, C, N,
splats_world, viewmats, intrins, dist_coeffs_buffer, image_width, image_height,
aabb, v_splats_screen, vr_splats_screen, h_splats_screen,
v_splats_world, vr_world_pos_buffer, h_world_pos_buffer,
vr_splats_world, h_splats_world, v_viewmats
);
}
@@ -0,0 +1,236 @@
#include "ProjectionFwd.cuh"
#include <gsplat/Utils.cuh>
#include <c10/cuda/CUDAStream.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
template<typename SplatPrimitive, gsplat::CameraModelType camera_model>
void projection_fused_fwd_kernel_wrapper(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
typename SplatPrimitive::Screen::Buffer splats_screen
);
template<typename SplatPrimitive>
inline std::tuple<
at::Tensor, // aabb
typename SplatPrimitive::Screen::TensorTupleProj // out splats
> launch_projection_fused_fwd_kernel(
const typename SplatPrimitive::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
typename SplatPrimitive::World::Tensor splats_world(in_splats);
uint32_t N = splats_world.size(); // number of gaussians
uint32_t C = viewmats.size(-3); // number of cameras
uint32_t B = splats_world.batchSize(); // number of batches
auto opt = splats_world.options();
at::Tensor aabb = at::empty({C, N, 4}, opt.dtype(at::kInt));
typename SplatPrimitive::Screen::Tensor splats_screen =
SplatPrimitive::Screen::Tensor::allocProjFwd(C, N, splats_world.options());
#define _LAUNCH_ARGS ( \
(cudaStream_t)at::cuda::getCurrentCUDAStream(), B, C, N, \
splats_world.buffer(), viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
image_width, image_height, near_plane, far_plane, \
(int4*)aabb.data_ptr<int32_t>(), splats_screen.buffer() \
)
if (camera_model == gsplat::CameraModelType::PINHOLE)
projection_fused_fwd_kernel_wrapper<SplatPrimitive, gsplat::CameraModelType::PINHOLE> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
projection_fused_fwd_kernel_wrapper<SplatPrimitive, gsplat::CameraModelType::FISHEYE> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
return std::make_tuple(aabb, splats_screen.tupleProjFwd());
}
// ================
// Vanilla3DGS
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
Vanilla3DGS::Screen::TensorTupleProj // out splats
> projection_3dgs_forward_tensor(
// inputs
const Vanilla3DGS::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<Vanilla3DGS>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
// ================
// MipSplatting
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
MipSplatting::Screen::TensorTupleProj // out splats
> projection_mip_forward_tensor(
// inputs
const MipSplatting::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<MipSplatting>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
// ================
// Vanilla3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
Vanilla3DGUT::Screen::TensorTupleProj // out splats
> projection_3dgut_forward_tensor(
// inputs
const Vanilla3DGUT::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<Vanilla3DGUT>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
// ================
// SphericalVoronoi3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj // out splats
> projection_3dgut_sv_forward_tensor(
// inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
int num_sv = std::get<5>(in_splats).size(-2);
#define _CASE(n) \
if (num_sv == n) return launch_projection_fused_fwd_kernel<SphericalVoronoi3DGUT<n>>( \
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs); \
_CASE(2) _CASE(3) _CASE(4) _CASE(5) _CASE(6) _CASE(7) _CASE(8)
#undef _CASE
throw std::invalid_argument("Unsupported num_sv");
}
// ================
// OpaqueTriangle
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
OpaqueTriangle::Screen::TensorTupleProj // out splats
> projection_opaque_triangle_forward_tensor(
// inputs
const OpaqueTriangle::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<OpaqueTriangle>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
// ================
// VoxelPrimitive
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
at::Tensor, // aabb
VoxelPrimitive::Screen::TensorTupleProj // out splats
> projection_voxel_forward_tensor(
// inputs
const VoxelPrimitive::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
return launch_projection_fused_fwd_kernel<VoxelPrimitive>(
in_splats, viewmats, intrins, image_width, image_height, near_plane, far_plane, camera_model, dist_coeffs);
}
+92 -112
View File
@@ -6,100 +6,27 @@
#include <ATen/Tensor.h>
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#include "types.cuh"
#include <c10/cuda/CUDAStream.h>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
#include "Primitive3DGS.cuh"
#include "Primitive3DGUT.cuh"
#include "Primitive3DGUT_SV.cuh"
#include "PrimitiveOpaqueTriangle.cuh"
#include "PrimitiveVoxel.cuh"
#include "common.cuh"
#include "types.cuh"
template<typename SplatPrimitive, gsplat::CameraModelType camera_model>
__global__ void projection_fused_fwd_kernel(
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
typename SplatPrimitive::Screen::Buffer splats_screen
) {
// parallelize over B * C * N.
uint32_t idx = cg::this_grid().thread_rank();
if (idx >= B * C * N) {
return;
}
const uint32_t bid = idx / (C * N); // batch id
const uint32_t cid = (idx / N) % C; // camera id
const uint32_t gid = idx % N; // gaussian id
// Load camera
viewmats += bid * C * 16 + cid * 16;
float4 intrin = intrins[bid * C + cid];
float3x3 R = {
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
typename SplatPrimitive::FwdProjCamera cam = {
R, t, fx, fy, cx, cy,
image_width, image_height,
near_plane, far_plane,
};
cam.dist_coeffs = dist_coeffs_buffer.load(bid * C + cid);
// Load splat
typename SplatPrimitive::World splat_world =
SplatPrimitive::World::load(splats_world, bid * N + gid);
// Projection
int4 aabb;
typename SplatPrimitive::Screen splat_screen;
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp(splat_world, cam, splat_screen, aabb);
break;
// case gsplat::CameraModelType::ORTHO: // orthographic projection
// SplatPrimitive::project_ortho(splat_world, cam, splat_screen, aabb);
// break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye(splat_world, cam, splat_screen, aabb);
break;
}
// Save results
aabb.x = min(max(aabb.x, 0), image_width-1);
aabb.y = min(max(aabb.y, 0), image_height-1);
aabb.z = min(max(aabb.z, 0), image_width-1);
aabb.w = min(max(aabb.w, 0), image_height-1);
if ((aabb.z-aabb.x)*(aabb.w-aabb.y) > 0) {
splat_screen.saveParamsToBuffer(splats_screen, idx);
aabbs[idx] = aabb;
} else {
aabbs[idx] = {0, 0, 0, 0};
}
}
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
template<typename SplatPrimitive>
inline std::tuple<
std::tuple<
at::Tensor, // aabb
typename SplatPrimitive::Screen::TensorTupleProj // out splats
> launch_projection_fused_fwd_kernel(
const typename SplatPrimitive::World::TensorTuple &in_splats,
Vanilla3DGS::Screen::TensorTupleProj // out splats
> projection_3dgs_forward_tensor(
// inputs
const Vanilla3DGS::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
@@ -108,36 +35,89 @@ inline std::tuple<
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
) {
typename SplatPrimitive::World::Tensor splats_world(in_splats);
uint32_t N = splats_world.size(); // number of gaussians
uint32_t C = viewmats.size(-3); // number of cameras
uint32_t B = splats_world.batchSize(); // number of batches
);
auto opt = splats_world.options();
at::Tensor aabb = at::empty({C, N, 4}, opt.dtype(at::kInt));
typename SplatPrimitive::Screen::Tensor splats_screen =
SplatPrimitive::Screen::Tensor::allocProjFwd(C, N, splats_world.options());
std::tuple<
at::Tensor, // aabb
MipSplatting::Screen::TensorTupleProj // out splats
> projection_mip_forward_tensor(
// inputs
const MipSplatting::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
#define _LAUNCH_ARGS \
<<<_LAUNCH_ARGS_1D(B*C*N, block)>>>( \
B, C, N, \
splats_world.buffer(), viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
image_width, image_height, near_plane, far_plane, \
(int4*)aabb.data_ptr<int32_t>(), splats_screen.buffer() \
)
constexpr uint block = 128;
if (camera_model == gsplat::CameraModelType::PINHOLE)
projection_fused_fwd_kernel<SplatPrimitive, gsplat::CameraModelType::PINHOLE> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
projection_fused_fwd_kernel<SplatPrimitive, gsplat::CameraModelType::FISHEYE> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
std::tuple<
at::Tensor, // aabb
Vanilla3DGUT::Screen::TensorTupleProj // out splats
> projection_3dgut_forward_tensor(
// inputs
const Vanilla3DGUT::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
#undef _LAUNCH_ARGS
return std::make_tuple(aabb, splats_screen.tupleProjFwd());
}
std::tuple<
at::Tensor, // aabb
SphericalVoronoi3DGUT_Default::Screen::TensorTupleProj // out splats
> projection_3dgut_sv_forward_tensor(
// inputs
const SphericalVoronoi3DGUT_Default::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
at::Tensor, // aabb
OpaqueTriangle::Screen::TensorTupleProj // out splats
> projection_opaque_triangle_forward_tensor(
// inputs
const OpaqueTriangle::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
std::tuple<
at::Tensor, // aabb
VoxelPrimitive::Screen::TensorTupleProj // out splats
> projection_voxel_forward_tensor(
// inputs
const VoxelPrimitive::World::TensorTuple &in_splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs
);
@@ -0,0 +1,120 @@
#include <cuda_runtime.h>
#include <cstdint>
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#ifndef NO_TORCH
#define NO_TORCH
#endif
#include "types.cuh"
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
template<typename SplatPrimitive, gsplat::CameraModelType camera_model>
__global__ void projection_fused_fwd_kernel(
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
typename SplatPrimitive::Screen::Buffer splats_screen
) {
// parallelize over B * C * N.
uint32_t idx = cg::this_grid().thread_rank();
if (idx >= B * C * N) {
return;
}
const uint32_t bid = idx / (C * N); // batch id
const uint32_t cid = (idx / N) % C; // camera id
const uint32_t gid = idx % N; // gaussian id
// Load camera
viewmats += bid * C * 16 + cid * 16;
float4 intrin = intrins[bid * C + cid];
float3x3 R = {
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
typename SplatPrimitive::FwdProjCamera cam = {
R, t, fx, fy, cx, cy,
image_width, image_height,
near_plane, far_plane,
};
cam.dist_coeffs = dist_coeffs_buffer.load(bid * C + cid);
// Load splat
typename SplatPrimitive::World splat_world =
SplatPrimitive::World::load(splats_world, bid * N + gid);
// Projection
int4 aabb;
typename SplatPrimitive::Screen splat_screen;
switch (camera_model) {
case gsplat::CameraModelType::PINHOLE: // perspective projection
SplatPrimitive::project_persp(splat_world, cam, splat_screen, aabb);
break;
// case gsplat::CameraModelType::ORTHO: // orthographic projection
// SplatPrimitive::project_ortho(splat_world, cam, splat_screen, aabb);
// break;
case gsplat::CameraModelType::FISHEYE: // fisheye projection
SplatPrimitive::project_fisheye(splat_world, cam, splat_screen, aabb);
break;
}
// Save results
aabb.x = min(max(aabb.x, 0), image_width-1);
aabb.y = min(max(aabb.y, 0), image_height-1);
aabb.z = min(max(aabb.z, 0), image_width-1);
aabb.w = min(max(aabb.w, 0), image_height-1);
if ((aabb.z-aabb.x)*(aabb.w-aabb.y) > 0) {
splat_screen.saveParamsToBuffer(splats_screen, idx);
aabbs[idx] = aabb;
} else {
aabbs[idx] = {0, 0, 0, 0};
}
}
template<typename SplatPrimitive, gsplat::CameraModelType camera_model>
void projection_fused_fwd_kernel_wrapper(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const typename SplatPrimitive::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
typename SplatPrimitive::Screen::Buffer splats_screen
) {
constexpr uint block = 128;
projection_fused_fwd_kernel<SplatPrimitive, gsplat::CameraModelType::PINHOLE>
<<<_CEIL_DIV(B*C*N, block), block, 0, stream>>>(
B, C, N,
splats_world, viewmats, intrins, dist_coeffs_buffer,
image_width, image_height, near_plane, far_plane,
aabbs, splats_screen
);
}
@@ -1,100 +0,0 @@
#include "RasterizationEval3DFwd.cuh"
// #include "RasterizationEval3DBwd_GSplat.cuh"
#include "RasterizationEval3DBwd.cuh"
#include "Primitive3DGUT.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_fwd(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
) {
if (output_distortion)
return rasterize_to_pixels_eval3d_fwd_tensor<Vanilla3DGUT, true, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
return rasterize_to_pixels_eval3d_fwd_tensor<Vanilla3DGUT, false, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_bwd(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
if (v_distortion_outputs.has_value())
return rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
@@ -1,58 +0,0 @@
// #include "RasterizationEval3DBwd_GSplat.cuh"
#include "RasterizationEval3DBwd.cuh"
#include "Primitive3DGUT.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor>, // v_viewmats
std::optional<Vanilla3DGUT::Screen::TensorTuple>, // jacobian residual product
std::optional<Vanilla3DGUT::Screen::TensorTuple> // hessian diagonal
> rasterize_to_pixels_3dgut_bwd_with_hessian_diagonal(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
if (v_distortion_outputs.has_value())
return _rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, true, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return _rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, false, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
@@ -1,99 +0,0 @@
#include "RasterizationEval3DFwd.cuh"
#include "RasterizationEval3DBwd.cuh"
#include "Primitive3DGUT_SV.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_sv_fwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
) {
if (output_distortion)
return rasterize_to_pixels_eval3d_fwd_tensor<SphericalVoronoi3DGUT_Default, true, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
return rasterize_to_pixels_eval3d_fwd_tensor<SphericalVoronoi3DGUT_Default, false, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
SphericalVoronoi3DGUT_Default::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_sv_bwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render_outputs,
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
if (v_distortion_outputs.has_value())
return rasterize_to_pixels_eval3d_bwd_tensor<SphericalVoronoi3DGUT_Default, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return rasterize_to_pixels_eval3d_bwd_tensor<SphericalVoronoi3DGUT_Default, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
@@ -20,156 +20,6 @@
std::tuple<
Vanilla3DGUT::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_fwd(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
);
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_bwd(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
);
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor>, // v_viewmats
std::optional<Vanilla3DGUT::Screen::TensorTuple>, // jacobian residual product
std::optional<Vanilla3DGUT::Screen::TensorTuple> // hessian diagonal
> rasterize_to_pixels_3dgut_bwd_with_hessian_diagonal(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
);
std::tuple<
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_sv_fwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
);
std::tuple<
SphericalVoronoi3DGUT_Default::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_sv_bwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render_outputs,
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
);
std::tuple<MipSplatting::RenderOutput::TensorTuple, at::Tensor, at::Tensor>
rasterize_to_pixels_mip_fwd(
// Gaussian parameters
@@ -233,63 +83,6 @@ std::tuple<
);
std::tuple<
VoxelPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_voxel_eval3d_fwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids // [n_isects]
);
std::tuple<
VoxelPrimitive::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_voxel_eval3d_bwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render_outputs,
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
VoxelPrimitive::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
);
std::tuple<
Vanilla3DGS::Screen::TensorTuple,
std::optional<Vanilla3DGS::Screen::TensorTuple>, // jacobian residual product
@@ -1,79 +0,0 @@
#include "RasterizationEval3DFwd.cuh"
#include "RasterizationEval3DBwd.cuh"
#include "PrimitiveVoxel.cuh"
/*[AutoHeaderGeneratorExport]*/
std::tuple<
VoxelPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_voxel_eval3d_fwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids // [n_isects]
) {
return rasterize_to_pixels_eval3d_fwd_tensor<VoxelPrimitive, true, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
VoxelPrimitive::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_voxel_eval3d_bwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render_outputs,
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
VoxelPrimitive::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
return rasterize_to_pixels_eval3d_bwd_tensor<VoxelPrimitive, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
@@ -0,0 +1,506 @@
#include "RasterizationEval3DBwd.cuh"
#include <gsplat/Utils.cuh>
#include "types.cuh"
#include "common.cuh"
#include <c10/cuda/CUDAStream.h>
template <
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
bool output_distortion,
bool output_viewmat_grad,
bool output_hessian_diagonal
>
void rasterize_to_pixels_eval3d_bwd_kernel_wrapper(
cudaStream_t stream,
const uint32_t I,
const uint32_t n_isects,
// fwd inputs
typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float *__restrict__ backgrounds, // [..., CDIM] or [nnz, CDIM]
const bool *__restrict__ masks, // [..., tile_height, tile_width]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [..., tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
// fwd outputs
const float *__restrict__ render_Ts, // [..., image_height, image_width, 1]
const int32_t *__restrict__ last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_output_buffer,
typename SplatPrimitive::RenderOutput::Buffer render2_output_buffer,
const float *__restrict__ loss_map_buffer, // [..., image_height, image_width, 1]
// grad outputs
typename SplatPrimitive::RenderOutput::Buffer v_render_output_buffer,
const float *__restrict__ v_render_alphas, // [..., image_height, image_width, 1]
typename SplatPrimitive::RenderOutput::Buffer v_distortions_output_buffer,
// grad inputs
typename SplatPrimitive::Screen::Buffer v_splat_buffer,
typename SplatPrimitive::Screen::Buffer vr_splat_buffer,
typename SplatPrimitive::Screen::Buffer h_splat_buffer,
float *__restrict__ v_viewmats // [B, C, 4, 4]
);
template <typename SplatPrimitive, bool output_distortion, bool output_hessian_diagonal>
inline void launch_rasterize_to_pixels_eval3d_bwd_kernel(
// Gaussian parameters
typename SplatPrimitive::Screen::Tensor splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., 3]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Tensor *render_outputs,
typename SplatPrimitive::RenderOutput::Tensor *render2_outputs,
const at::Tensor *loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
typename SplatPrimitive::RenderOutput::Tensor v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
typename SplatPrimitive::RenderOutput::Tensor *v_distortion_outputs,
// outputs
typename SplatPrimitive::Screen::Tensor v_splats,
typename SplatPrimitive::Screen::Tensor *vr_splats,
typename SplatPrimitive::Screen::Tensor *h_splats,
std::optional<at::Tensor> v_viewmats
) {
bool packed = splats.isPacked();
uint32_t N = packed ? 0 : splats.size(); // number of gaussians
uint32_t I = render_Ts.numel() / (image_height * image_width); // number of images
uint32_t tile_height = tile_offsets.size(-2);
uint32_t tile_width = tile_offsets.size(-1);
uint32_t n_isects = flatten_ids.size(0);
if (n_isects == 0) {
// skip the kernel launch if there are no elements
return;
}
typename SplatPrimitive::Screen::Buffer vr_splats_buffer;
typename SplatPrimitive::Screen::Buffer h_splats_buffer;
if (output_hessian_diagonal) {
vr_splats_buffer = vr_splats->buffer();
h_splats_buffer = h_splats->buffer();
}
#define _LAUNCH_ARGS ( \
(cudaStream_t)at::cuda::getCurrentCUDAStream(), I, n_isects, \
splats.buffer(), \
viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
backgrounds.has_value() ? backgrounds.value().data_ptr<float>() : nullptr, \
masks.has_value() ? masks.value().data_ptr<bool>() : nullptr, \
image_width, image_height, tile_width, tile_height, \
tile_offsets.data_ptr<int32_t>(), flatten_ids.data_ptr<int32_t>(), \
render_Ts.data_ptr<float>(), last_ids.data_ptr<int32_t>(), \
output_distortion ? render_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
output_distortion ? render2_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
output_hessian_diagonal ? loss_map->data_ptr<float>() : nullptr, \
v_render_outputs.buffer(), v_render_alphas.data_ptr<float>(), \
output_distortion ? v_distortion_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
v_splats.buffer(), vr_splats_buffer, h_splats_buffer, \
v_viewmats.has_value() ? v_viewmats.value().data_ptr<float>() : nullptr \
)
if (camera_model == gsplat::CameraModelType::PINHOLE) {
if (v_viewmats.has_value())
rasterize_to_pixels_eval3d_bwd_kernel_wrapper<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion, true, output_hessian_diagonal> _LAUNCH_ARGS;
else
rasterize_to_pixels_eval3d_bwd_kernel_wrapper<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion, false, output_hessian_diagonal> _LAUNCH_ARGS;
}
else if (camera_model == gsplat::CameraModelType::FISHEYE) {
if (v_viewmats.has_value())
rasterize_to_pixels_eval3d_bwd_kernel_wrapper<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion, true, output_hessian_diagonal> _LAUNCH_ARGS;
else
rasterize_to_pixels_eval3d_bwd_kernel_wrapper<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion, false, output_hessian_diagonal> _LAUNCH_ARGS;
}
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
}
template<typename SplatPrimitive, bool output_distortion, bool output_hessian_diagonal>
inline std::tuple<
typename SplatPrimitive::Screen::TensorTuple,
std::optional<at::Tensor>, // v_viewmats
std::optional<typename SplatPrimitive::Screen::TensorTuple>, // jacobian residual product
std::optional<typename SplatPrimitive::Screen::TensorTuple> // hessian diagonal
> _rasterize_to_pixels_eval3d_bwd_tensor(
// Gaussian parameters
typename SplatPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> render_outputs_tuple,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> render2_outputs_tuple,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
typename SplatPrimitive::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> v_distortion_outputs_tuple,
bool need_viewmat_grad
) {
DEVICE_GUARD(tile_offsets);
CHECK_INPUT(tile_offsets);
CHECK_INPUT(flatten_ids);
CHECK_INPUT(viewmats);
CHECK_INPUT(intrins);
CHECK_INPUT(render_Ts);
CHECK_INPUT(last_ids);
CHECK_INPUT(v_render_alphas);
if (backgrounds.has_value())
CHECK_INPUT(backgrounds.value());
if (masks.has_value())
CHECK_INPUT(masks.value());
if (loss_map.has_value())
CHECK_INPUT(loss_map.value());
typename SplatPrimitive::Screen::Tensor splats(splats_tuple);
typename SplatPrimitive::Screen::Tensor v_splats = splats.allocRasterBwd();
std::optional<at::Tensor> v_viewmats = need_viewmat_grad ?
(std::optional<at::Tensor>)zeros_like<float>(viewmats) : (std::optional<at::Tensor>)std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> render_outputs = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> render2_outputs = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> v_distortion_outputs = std::nullopt;
if (output_distortion) {
render_outputs = render_outputs_tuple;
render2_outputs = render2_outputs_tuple;
v_distortion_outputs = v_distortion_outputs_tuple;
}
std::optional<typename SplatPrimitive::Screen::Tensor> vr_splats = std::nullopt;
std::optional<typename SplatPrimitive::Screen::Tensor> h_splats = std::nullopt;
if (output_hessian_diagonal) {
vr_splats = splats.allocRasterBwd();
h_splats = splats.allocRasterBwd();
}
launch_rasterize_to_pixels_eval3d_bwd_kernel<SplatPrimitive, output_distortion, output_hessian_diagonal>(
splats,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids,
output_distortion ? &render_outputs.value() : nullptr,
output_distortion ? &render2_outputs.value() : nullptr,
output_hessian_diagonal ? &loss_map.value() : nullptr,
v_render_outputs, v_render_alphas,
output_distortion ? &v_distortion_outputs.value() : nullptr,
v_splats,
output_hessian_diagonal ? &vr_splats.value() : nullptr,
output_hessian_diagonal ? &h_splats.value() : nullptr,
v_viewmats
);
if (output_hessian_diagonal)
return std::make_tuple(v_splats.tupleRasterBwd(), v_viewmats,
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)vr_splats.value().tupleRasterBwd(),
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)h_splats.value().tupleRasterBwd());
return std::make_tuple(v_splats.tupleRasterBwd(), v_viewmats,
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)std::nullopt,
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)std::nullopt);
}
template<typename SplatPrimitive, bool output_distortion>
inline std::tuple<
typename SplatPrimitive::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_eval3d_bwd_tensor(
// Gaussian parameters
typename SplatPrimitive::Screen::TensorTuple &splats_tuple,
const at::Tensor &viewmats, // [..., C, 4, 4]
const at::Tensor &intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType &camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> &backgrounds, // [..., channels]
const std::optional<at::Tensor> &masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor &tile_offsets, // [..., tile_height, tile_width]
const at::Tensor &flatten_ids, // [n_isects]
// forward outputs
const at::Tensor &render_Ts, // [..., image_height, image_width, 1]
const at::Tensor &last_ids, // [..., image_height, image_width]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> &render_outputs,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> &render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
typename SplatPrimitive::RenderOutput::TensorTuple &v_render_outputs,
const at::Tensor &v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> &v_distortion_outputs,
bool need_viewmat_grad
) {
auto [v_splats, v_viewmat, vr_splats, h_splats] =
_rasterize_to_pixels_eval3d_bwd_tensor<SplatPrimitive, output_distortion, false>
(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return std::make_tuple(v_splats, v_viewmat);
}
// ================
// Vanilla3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_bwd(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
if (v_distortion_outputs.has_value())
return rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor>, // v_viewmats
std::optional<Vanilla3DGUT::Screen::TensorTuple>, // jacobian residual product
std::optional<Vanilla3DGUT::Screen::TensorTuple> // hessian diagonal
> rasterize_to_pixels_3dgut_bwd_with_hessian_diagonal(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
if (v_distortion_outputs.has_value())
return _rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, true, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return _rasterize_to_pixels_eval3d_bwd_tensor<Vanilla3DGUT, false, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
// ================
// SphericalVoronoi3DGUT
// ================
// TODO: Is this the same as Vanilla3DGUT?
/*[AutoHeaderGeneratorExport]*/
std::tuple<
SphericalVoronoi3DGUT_Default::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_sv_bwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render_outputs,
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
if (v_distortion_outputs.has_value())
return rasterize_to_pixels_eval3d_bwd_tensor<SphericalVoronoi3DGUT_Default, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return rasterize_to_pixels_eval3d_bwd_tensor<SphericalVoronoi3DGUT_Default, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
// ================
// VoxelPrimitive
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
VoxelPrimitive::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_voxel_eval3d_bwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render_outputs,
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
VoxelPrimitive::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
return rasterize_to_pixels_eval3d_bwd_tensor<VoxelPrimitive, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
}
@@ -1,507 +1,31 @@
#pragma once
// Modified from https://github.com/nerfstudio-project/gsplat/blob/main/gsplat/cuda/csrc/RasterizeToPixels3DGSBwd.cu
#include <cuda.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <ATen/Tensor.h>
#ifdef __CUDACC__
#include "generated/slang.cuh"
namespace SlangProjectionUtils {
#include "generated/set_namespace.cuh"
#include "generated/projection_utils.cuh"
}
#endif
// #include "Primitive3DGS.cuh"
#include "Primitive3DGUT.cuh"
#include "Primitive3DGUT_SV.cuh"
// #include "PrimitiveOpaqueTriangle.cuh"
#include "PrimitiveVoxel.cuh"
#include "types.cuh"
#include "common.cuh"
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#include <cub/cub.cuh>
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
constexpr uint SPLAT_BATCH_SIZE_NO_DISTORTION = 128;
constexpr uint SPLAT_BATCH_SIZE_WITH_DISTORTION = WARP_SIZE;
template <
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
bool output_distortion,
bool output_viewmat_grad,
bool output_hessian_diagonal
>
__global__ void rasterize_to_pixels_eval3d_bwd_kernel(
const uint32_t I,
const uint32_t n_isects,
// fwd inputs
typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float *__restrict__ backgrounds, // [..., CDIM] or [nnz, CDIM]
const bool *__restrict__ masks, // [..., tile_height, tile_width]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [..., tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
// fwd outputs
const float *__restrict__ render_Ts, // [..., image_height, image_width, 1]
const int32_t *__restrict__ last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_output_buffer,
typename SplatPrimitive::RenderOutput::Buffer render2_output_buffer,
const float *__restrict__ loss_map_buffer, // [..., image_height, image_width, 1]
// grad outputs
typename SplatPrimitive::RenderOutput::Buffer v_render_output_buffer,
const float *__restrict__ v_render_alphas, // [..., image_height, image_width, 1]
typename SplatPrimitive::RenderOutput::Buffer v_distortions_output_buffer,
// grad inputs
typename SplatPrimitive::Screen::Buffer v_splat_buffer,
typename SplatPrimitive::Screen::Buffer vr_splat_buffer,
typename SplatPrimitive::Screen::Buffer h_splat_buffer,
float *__restrict__ v_viewmats // [B, C, 4, 4]
) {
auto block = cg::this_thread_block();
cg::thread_block_tile<WARP_SIZE> warp = cg::tiled_partition<WARP_SIZE>(block);
uint32_t image_id = block.group_index().x;
uint32_t tile_id = block.group_index().y * tile_width + block.group_index().z;
uint32_t thread_id = block.thread_rank();
tile_offsets += image_id * tile_height * tile_width;
render_Ts += image_id * image_height * image_width;
last_ids += image_id * image_height * image_width;
v_render_alphas += image_id * image_height * image_width;
if (backgrounds != nullptr) {
backgrounds += image_id;
}
if (masks != nullptr) {
masks += image_id * tile_height * tile_width;
}
// when the mask is provided, do nothing and return if
// this tile is labeled as False
if (masks != nullptr && !masks[tile_id]) {
return;
}
// Load camera
viewmats += image_id * 16; // world to camera
float4 intrin = intrins[image_id];
float3x3 R = { // row major
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
CameraDistortionCoeffs dist_coeffs = dist_coeffs_buffer.load(image_id);
constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE;
// load pixels
__shared__ float4 shared_ray_d_pix_bin_final[BLOCK_SIZE];
__shared__ float2 pix_Ts_with_grad[BLOCK_SIZE];
__shared__ typename SplatPrimitive::RenderOutput v_pix_colors[BLOCK_SIZE];
// __shared__ float pix_background[CDIM]; // TODO
__shared__ typename SplatPrimitive::RenderOutput pix_colors[output_distortion ? BLOCK_SIZE : 1];
__shared__ typename SplatPrimitive::RenderOutput pix2_colors[output_distortion ? BLOCK_SIZE : 1];
__shared__ typename SplatPrimitive::RenderOutput v_distortion_out[output_distortion ? BLOCK_SIZE : 1];
__shared__ float hess_weight_map[output_hessian_diagonal ? BLOCK_SIZE : 1];
float3 ray_o = SlangProjectionUtils::transform_ray_o(R, t);
float3 total_v_ray_o = make_float3(0.0f, 0.0f, 0.0f);
__shared__ float3 shared_v_ray_d[output_viewmat_grad ? BLOCK_SIZE : 1];
constexpr uint SPLAT_BATCH_SIZE_CONST = output_distortion ?
SPLAT_BATCH_SIZE_WITH_DISTORTION : SPLAT_BATCH_SIZE_NO_DISTORTION;
#pragma unroll
for (uint pix_id0 = 0; pix_id0 < BLOCK_SIZE; pix_id0 += SPLAT_BATCH_SIZE_CONST) {
static_assert(BLOCK_SIZE % SPLAT_BATCH_SIZE_CONST == 0);
uint pix_id_local = pix_id0 + thread_id;
int pix_x = block.group_index().z * TILE_SIZE + pix_id_local % TILE_SIZE;
int pix_y = block.group_index().y * TILE_SIZE + pix_id_local / TILE_SIZE;
uint pix_id_global = pix_y * image_width + pix_x;
uint pix_id_image_global = image_id * image_height * image_width + pix_id_global;
bool inside = (pix_x < image_width && pix_y < image_height);
int32_t bin_final = (inside ? last_ids[pix_id_global] : 0);
pix_Ts_with_grad[pix_id_local] = {
(inside ? render_Ts[pix_id_global] : 0.0f),
(inside ? -v_render_alphas[pix_id_global] : 0.0f)
};
v_pix_colors[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(v_render_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
const float px = (float)pix_x + 0.5f;
const float py = (float)pix_y + 0.5f;
float3 raydir;
inside &= SlangProjectionUtils::generate_ray(
{(px-cx)/fx, (py-cy)/fy},
camera_model == gsplat::CameraModelType::FISHEYE, dist_coeffs,
&raydir
);
float3 ray_d = SlangProjectionUtils::transform_ray_d(R, raydir); // mul(raydir, R);
shared_ray_d_pix_bin_final[pix_id_local] =
{ray_d.x, ray_d.y, ray_d.z, __int_as_float(bin_final)};
if (output_distortion) {
pix_colors[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(render_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
pix2_colors[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(render2_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
v_distortion_out[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(v_distortions_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
}
if (output_viewmat_grad) {
shared_v_ray_d[pix_id_local] = make_float3(0.0f, 0.0f, 0.0f);
}
if (output_hessian_diagonal) {
// https://www.desmos.com/calculator/ld9wg7cuxz
hess_weight_map[pix_id_local] = (loss_map_buffer != nullptr && inside) ?
0.5f / fmaxf(loss_map_buffer[pix_id_image_global], 1e-6f) : 0.0f;
}
}
block.sync();
// threads fist load splats, then swept through pixels
// do this in batches
int32_t range_start = tile_offsets[tile_id];
int32_t range_end =
(image_id == I - 1) && (tile_id == tile_width * tile_height - 1)
? n_isects
: tile_offsets[tile_id + 1];
uint SPLAT_BATCH_SIZE = SPLAT_BATCH_SIZE_CONST;
if (SPLAT_BATCH_SIZE_CONST > WARP_SIZE) {
SPLAT_BATCH_SIZE = (uint)sqrtf((float)(range_end - range_start) * (float)BLOCK_SIZE);
// SPLAT_BATCH_SIZE = min(SPLAT_BATCH_SIZE_CONST, (SPLAT_BATCH_SIZE + WARP_SIZE) & ~(WARP_SIZE-1));
SPLAT_BATCH_SIZE = min(SPLAT_BATCH_SIZE_CONST, max(SPLAT_BATCH_SIZE, 1u));
}
const uint32_t num_splat_batches =
_CEIL_DIV(range_end - range_start, SPLAT_BATCH_SIZE);
// if (warp.thread_rank() == 0)
// printf("range_start=%d range_end=%d num_splat_batches=%u\n", range_start, range_end, num_splat_batches);
for (uint32_t splat_b = 0; splat_b < num_splat_batches; ++splat_b) {
const int32_t splat_batch_end = range_end - 1 - SPLAT_BATCH_SIZE * splat_b;
const int32_t splat_batch_size = min(SPLAT_BATCH_SIZE, splat_batch_end + 1 - range_start);
const int32_t splat_idx = splat_batch_end - thread_id;
// load splats
typename SplatPrimitive::Screen splat;
uint32_t splat_gid;
if (splat_idx >= range_start) {
splat_gid = flatten_ids[splat_idx]; // flatten index in [I * N] or [nnz]
splat = SplatPrimitive::Screen::loadWithPrecompute(splat_buffer, splat_gid);
}
// accumulate gradient
typename SplatPrimitive::Screen v_splat = SplatPrimitive::Screen::zero();
typename SplatPrimitive::Screen vr_splat = SplatPrimitive::Screen::zero();
typename SplatPrimitive::Screen h_splat = SplatPrimitive::Screen::zero();
// thread 0 takes last splat, 1 takes second last, etc.
// at t=0, thread 0 (splat -1) undo pixel 0
// at t=1, thread 0 (splat -1) undo pixel 1, thread 1 (splat -2) undo pixel 0
// ......
// process gaussians in the current batch for this pixel
// 0 index is the furthest back gaussian in the batch
for (int t = 0; t < splat_batch_size + BLOCK_SIZE - 1; ++t, __syncwarp()) {
int pix_id = t - thread_id;
if (pix_id < 0 || pix_id >= BLOCK_SIZE || splat_idx < range_start)
continue;
float4 ray_d_pix_bin_final = shared_ray_d_pix_bin_final[pix_id];
if (splat_idx > __float_as_int(ray_d_pix_bin_final.w))
continue;
// evaluate alpha and early skip
float3 ray_d = {ray_d_pix_bin_final.x, ray_d_pix_bin_final.y, ray_d_pix_bin_final.z};
float alpha = splat.evaluate_alpha(ray_o, ray_d);
if (alpha <= ALPHA_THRESHOLD || dot(ray_d, ray_d) == 0.0f)
continue;
// printf("t=%d, thread %u, splat %d (%u), pix_id %d, pix %d %d\n", t, thread_id, splat_idx-range_start, splat_gid, pix_id, pix_global_x, pix_global_y);
// forward:
// \left(c_{1},T_{1}\right)=\left(c_{0}+\alpha_{i}T_{0}c_{i},\ T_{0}\left(1-\alpha_{i}\right)\right)
float T1 = pix_Ts_with_grad[pix_id].x;
float v_T1 = pix_Ts_with_grad[pix_id].y;
// undo pixel:
// T_{0}=\frac{T_{1}}{1-\alpha_{i}}
float ra = 1.0f / (1.0f - alpha);
float T0 = T1 * ra;
typename SplatPrimitive::RenderOutput color = splat.evaluate_color(ray_o, ray_d);
typename SplatPrimitive::RenderOutput v_c = v_pix_colors[pix_id];
// gradient to alpha:
// \frac{dL}{d\alpha_{i}}
// = \frac{dL}{dc_{1}}\frac{dc_{1}}{d\alpha_{i}}+\frac{dL}{dT_{1}}\frac{dT_{1}}{d\alpha_{i}}
// = T_{0}\frac{dL}{dc_{1}}c_{i}-\frac{dL}{dT_{1}}T_{0}
float v_alpha = T0 * color.dot(v_c) -v_T1 * T0;
// gradient to color:
// \frac{dL}{dc_{i}}
// = \frac{dL}{dc_{1}}\frac{dc_{1}}{dc_{i}}
// = \alpha_{i}T_{0}\frac{dL}{dc_{1}}
typename SplatPrimitive::RenderOutput v_color = v_c * (alpha * T0);
// update pixel gradient:
// \frac{dL}{dT_{0}}
// = \frac{dL}{dc_{1}}\frac{dc_{1}}{dT_{0}}+\frac{dL}{dT_{1}}\frac{dT_{1}}{dT_{0}}
// = \alpha_{i}\frac{dL}{dc_{1}}c_{i}+\frac{dL}{dT_{1}}\left(1-\alpha_{i}\right)
float v_T0 = alpha * color.dot(v_c) + v_T1 * (1.0f - alpha);
// distortion
if (output_distortion) {
// \left(d_{1},s_{1}\right)=\left(d_{0}+\alpha_{i}T_{0}\left(c_{i}^{2}\left(1-T_{0}\right)-2c_{i}c_{0}+s_{0}\right),\ s_{0}+\alpha_{i}T_{0}c_{i}^{2}\right)
// \frac{dL}{ds}=0
typename SplatPrimitive::RenderOutput v_dist = v_distortion_out[pix_id];
typename SplatPrimitive::RenderOutput c0 =
pix_colors[pix_id] + color * -alpha * T0;
typename SplatPrimitive::RenderOutput s0 =
pix2_colors[pix_id] + color * color * -alpha * T0;
// \frac{dL}{d\alpha_{i}}=\frac{dL}{dd_{1}}\frac{dd_{1}}{d\alpha_{i}}=T_{0}\left(c_{i}^{2}\left(1-T_{0}\right)-2c_{i}c_{0}+s_{0}\right)\frac{dL}{dd_{1}}
v_alpha += T0 * (
color * color * (1.0f-T0) +
color * c0 * -2.0f + s0
).dot(v_dist);
// \frac{dL}{dc_{i}}=\frac{dL}{dd_{1}}\frac{dd_{1}}{dc_{i}}=2\alpha_{i}T_{0}\left(c_{i}\left(1-T_{0}\right)-c_{0}\right)\frac{dL}{dd_{1}}
v_color += (
color * (1.0f-T0) +
c0 * -1.0f
) * v_dist * (2.0f * alpha * T0);
// \alpha_{i}\left(c_{i}^{2}\left(1-2T_{0}\right)-2c_{i}c_{0}+s_{0}\right)\frac{dL}{dd_{1}}
v_T0 += alpha * (
color * color * (1.0f-2.0f*T0) +
color * c0 * -2.0f + s0
).dot(v_dist);
// undo pixel state
pix_colors[pix_id] = c0;
pix2_colors[pix_id] = s0;
}
// backward diff splat
float3 v_ray_o_alpha, v_ray_d_alpha;
float3 v_ray_o_color, v_ray_d_color;
if (output_hessian_diagonal) {
typename SplatPrimitive::Screen v_splat_temp = SplatPrimitive::Screen::zero();
v_splat_temp.addGradient(splat.evaluate_alpha_vjp(ray_o, ray_d, v_alpha, v_ray_o_alpha, v_ray_d_alpha));
v_splat_temp.addGradient(splat.evaluate_color_vjp(ray_o, ray_d, v_color, v_ray_o_color, v_ray_d_color));
v_splat.addGradient(v_splat_temp);
// https://www.desmos.com/calculator/ld9wg7cuxz
float weight = 1.0f / sqrtf(fmaxf(v_c.dot(v_c) / 3.0f, 1e-30f));
vr_splat.addGradient(v_splat_temp, weight);
weight = hess_weight_map[pix_id] * weight * weight;
splat.precomputeBackward(v_splat_temp);
h_splat.addGaussNewtonHessianDiagonal(v_splat_temp, weight);
} else {
v_splat.addGradient(splat.evaluate_alpha_vjp(ray_o, ray_d, v_alpha, v_ray_o_alpha, v_ray_d_alpha));
v_splat.addGradient(splat.evaluate_color_vjp(ray_o, ray_d, v_color, v_ray_o_color, v_ray_d_color));
}
if (output_viewmat_grad) {
total_v_ray_o += v_ray_o_alpha + v_ray_o_color;
shared_v_ray_d[pix_id] += v_ray_d_alpha + v_ray_d_color;
}
// update pixel states
pix_Ts_with_grad[pix_id] = { T0, v_T0 };
// v_pix_colors remains the same
}
// accumulate gradient
if (splat_idx >= range_start) {
splat.precomputeBackward(v_splat);
v_splat.atomicAddToBuffer(v_splat_buffer, splat_gid);
if (output_hessian_diagonal) {
splat.precomputeBackward(vr_splat);
vr_splat.atomicAddToBuffer(vr_splat_buffer, splat_gid);
h_splat.atomicAddToBuffer(h_splat_buffer, splat_gid);
}
}
}
if (output_viewmat_grad) {
// accumulate to viewmat gradient
float3x3 v_R;
float3 v_t;
// gradient from ray_o (will fill v_R and v_t)
SlangProjectionUtils::transform_ray_o_vjp(R, t, total_v_ray_o, &v_R, &v_t);
// gradient from ray_d
#pragma unroll
for (uint pix_id0 = 0; pix_id0 < BLOCK_SIZE; pix_id0 += SPLAT_BATCH_SIZE_CONST) {
uint pix_id_local = pix_id0 + thread_id;
float4 ray_d_pix_bin_final = shared_ray_d_pix_bin_final[pix_id_local];
float3 raydir = SlangProjectionUtils::undo_transform_ray_d(R,
make_float3(
ray_d_pix_bin_final.x,
ray_d_pix_bin_final.y,
ray_d_pix_bin_final.z
)
);
float3 v_ray_d = shared_v_ray_d[pix_id_local];
float3x3 v_R_delta;
float3 temp;
SlangProjectionUtils::transform_ray_d_vjp(R, raydir, v_ray_d, &v_R_delta, &temp);
v_R = v_R + v_R_delta;
}
// atomic add to global viewmat gradient
if (v_viewmats != nullptr) {
float *v_viewmat = v_viewmats + image_id * 16;
float temp;
#define _ATOMIC_ADD(ptr, offset, value) do { \
temp = isfinite(value) ? value : 0.0f; \
warpSum(temp, warp); \
if (warp.thread_rank() == 0 && temp != 0.0f) \
atomicAdd((ptr) + (offset), (temp)); \
} while(0)
_ATOMIC_ADD(v_viewmat, 0, v_R[0].x);
_ATOMIC_ADD(v_viewmat, 1, v_R[0].y);
_ATOMIC_ADD(v_viewmat, 2, v_R[0].z);
_ATOMIC_ADD(v_viewmat, 3, v_t.x);
_ATOMIC_ADD(v_viewmat, 4, v_R[1].x);
_ATOMIC_ADD(v_viewmat, 5, v_R[1].y);
_ATOMIC_ADD(v_viewmat, 6, v_R[1].z);
_ATOMIC_ADD(v_viewmat, 7, v_t.y);
_ATOMIC_ADD(v_viewmat, 8, v_R[2].x);
_ATOMIC_ADD(v_viewmat, 9, v_R[2].y);
_ATOMIC_ADD(v_viewmat, 10, v_R[2].z);
_ATOMIC_ADD(v_viewmat, 11, v_t.z);
#undef _ATOMIC_ADD
}
}
}
template <typename SplatPrimitive, bool output_distortion, bool output_hessian_diagonal>
inline void launch_rasterize_to_pixels_eval3d_bwd_kernel(
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_bwd(
// Gaussian parameters
typename SplatPrimitive::Screen::Tensor splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., 3]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Tensor *render_outputs,
typename SplatPrimitive::RenderOutput::Tensor *render2_outputs,
const at::Tensor *loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
typename SplatPrimitive::RenderOutput::Tensor v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
typename SplatPrimitive::RenderOutput::Tensor *v_distortion_outputs,
// outputs
typename SplatPrimitive::Screen::Tensor v_splats,
typename SplatPrimitive::Screen::Tensor *vr_splats,
typename SplatPrimitive::Screen::Tensor *h_splats,
std::optional<at::Tensor> v_viewmats
) {
bool packed = splats.isPacked();
uint32_t N = packed ? 0 : splats.size(); // number of gaussians
uint32_t I = render_Ts.numel() / (image_height * image_width); // number of images
uint32_t tile_height = tile_offsets.size(-2);
uint32_t tile_width = tile_offsets.size(-1);
uint32_t n_isects = flatten_ids.size(0);
// Each block covers a tile on the image. In total there are
// I * tile_height * tile_width blocks.
dim3 threads = {output_distortion ?
SPLAT_BATCH_SIZE_WITH_DISTORTION : SPLAT_BATCH_SIZE_NO_DISTORTION,
1, 1};
dim3 grid = {I, tile_height, tile_width};
if (n_isects == 0) {
// skip the kernel launch if there are no elements
return;
}
typename SplatPrimitive::Screen::Buffer vr_splats_buffer;
typename SplatPrimitive::Screen::Buffer h_splats_buffer;
if (output_hessian_diagonal) {
vr_splats_buffer = vr_splats->buffer();
h_splats_buffer = h_splats->buffer();
}
#define _LAUNCH_ARGS <<<grid, threads, 0, at::cuda::getCurrentCUDAStream()>>>( \
I, n_isects, \
splats.buffer(), \
viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
backgrounds.has_value() ? backgrounds.value().data_ptr<float>() : nullptr, \
masks.has_value() ? masks.value().data_ptr<bool>() : nullptr, \
image_width, image_height, tile_width, tile_height, \
tile_offsets.data_ptr<int32_t>(), flatten_ids.data_ptr<int32_t>(), \
render_Ts.data_ptr<float>(), last_ids.data_ptr<int32_t>(), \
output_distortion ? render_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
output_distortion ? render2_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
output_hessian_diagonal ? loss_map->data_ptr<float>() : nullptr, \
v_render_outputs.buffer(), v_render_alphas.data_ptr<float>(), \
output_distortion ? v_distortion_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
v_splats.buffer(), vr_splats_buffer, h_splats_buffer, \
v_viewmats.has_value() ? v_viewmats.value().data_ptr<float>() : nullptr \
)
if (camera_model == gsplat::CameraModelType::PINHOLE) {
if (v_viewmats.has_value())
rasterize_to_pixels_eval3d_bwd_kernel<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion, true, output_hessian_diagonal> _LAUNCH_ARGS;
else
rasterize_to_pixels_eval3d_bwd_kernel<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion, false, output_hessian_diagonal> _LAUNCH_ARGS;
}
else if (camera_model == gsplat::CameraModelType::FISHEYE) {
if (v_viewmats.has_value())
rasterize_to_pixels_eval3d_bwd_kernel<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion, true, output_hessian_diagonal> _LAUNCH_ARGS;
else
rasterize_to_pixels_eval3d_bwd_kernel<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion, false, output_hessian_diagonal> _LAUNCH_ARGS;
}
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
}
template<typename SplatPrimitive, bool output_distortion, bool output_hessian_diagonal>
inline std::tuple<
typename SplatPrimitive::Screen::TensorTuple,
std::optional<at::Tensor>, // v_viewmats
std::optional<typename SplatPrimitive::Screen::TensorTuple>, // jacobian residual product
std::optional<typename SplatPrimitive::Screen::TensorTuple> // hessian diagonal
> _rasterize_to_pixels_eval3d_bwd_tensor(
// Gaussian parameters
typename SplatPrimitive::Screen::TensorTuple splats_tuple,
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
@@ -517,119 +41,110 @@ inline std::tuple<
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> render_outputs_tuple,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> render2_outputs_tuple,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
typename SplatPrimitive::RenderOutput::TensorTuple v_render_outputs,
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> v_distortion_outputs_tuple,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
DEVICE_GUARD(tile_offsets);
CHECK_INPUT(tile_offsets);
CHECK_INPUT(flatten_ids);
CHECK_INPUT(viewmats);
CHECK_INPUT(intrins);
CHECK_INPUT(render_Ts);
CHECK_INPUT(last_ids);
CHECK_INPUT(v_render_alphas);
if (backgrounds.has_value())
CHECK_INPUT(backgrounds.value());
if (masks.has_value())
CHECK_INPUT(masks.value());
if (loss_map.has_value())
CHECK_INPUT(loss_map.value());
typename SplatPrimitive::Screen::Tensor splats(splats_tuple);
typename SplatPrimitive::Screen::Tensor v_splats = splats.allocRasterBwd();
std::optional<at::Tensor> v_viewmats = need_viewmat_grad ?
(std::optional<at::Tensor>)zeros_like<float>(viewmats) : (std::optional<at::Tensor>)std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> render_outputs = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> render2_outputs = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> v_distortion_outputs = std::nullopt;
if (output_distortion) {
render_outputs = render_outputs_tuple;
render2_outputs = render2_outputs_tuple;
v_distortion_outputs = v_distortion_outputs_tuple;
}
std::optional<typename SplatPrimitive::Screen::Tensor> vr_splats = std::nullopt;
std::optional<typename SplatPrimitive::Screen::Tensor> h_splats = std::nullopt;
if (output_hessian_diagonal) {
vr_splats = splats.allocRasterBwd();
h_splats = splats.allocRasterBwd();
}
launch_rasterize_to_pixels_eval3d_bwd_kernel<SplatPrimitive, output_distortion, output_hessian_diagonal>(
splats,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids,
output_distortion ? &render_outputs.value() : nullptr,
output_distortion ? &render2_outputs.value() : nullptr,
output_hessian_diagonal ? &loss_map.value() : nullptr,
v_render_outputs, v_render_alphas,
output_distortion ? &v_distortion_outputs.value() : nullptr,
v_splats,
output_hessian_diagonal ? &vr_splats.value() : nullptr,
output_hessian_diagonal ? &h_splats.value() : nullptr,
v_viewmats
);
if (output_hessian_diagonal)
return std::make_tuple(v_splats.tupleRasterBwd(), v_viewmats,
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)vr_splats.value().tupleRasterBwd(),
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)h_splats.value().tupleRasterBwd());
return std::make_tuple(v_splats.tupleRasterBwd(), v_viewmats,
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)std::nullopt,
(std::optional<typename SplatPrimitive::Screen::TensorTuple>)std::nullopt);
}
);
template<typename SplatPrimitive, bool output_distortion>
inline std::tuple<
typename SplatPrimitive::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_eval3d_bwd_tensor(
std::tuple<
Vanilla3DGUT::Screen::TensorTuple,
std::optional<at::Tensor>, // v_viewmats
std::optional<Vanilla3DGUT::Screen::TensorTuple>, // jacobian residual product
std::optional<Vanilla3DGUT::Screen::TensorTuple> // hessian diagonal
> rasterize_to_pixels_3dgut_bwd_with_hessian_diagonal(
// Gaussian parameters
typename SplatPrimitive::Screen::TensorTuple &splats_tuple,
const at::Tensor &viewmats, // [..., C, 4, 4]
const at::Tensor &intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType &camera_model,
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> &backgrounds, // [..., channels]
const std::optional<at::Tensor> &masks, // [..., tile_height, tile_width]
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor &tile_offsets, // [..., tile_height, tile_width]
const at::Tensor &flatten_ids, // [n_isects]
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor &render_Ts, // [..., image_height, image_width, 1]
const at::Tensor &last_ids, // [..., image_height, image_width]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> &render_outputs,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> &render2_outputs,
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render_outputs,
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
typename SplatPrimitive::RenderOutput::TensorTuple &v_render_outputs,
const at::Tensor &v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> &v_distortion_outputs,
Vanilla3DGUT::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename Vanilla3DGUT::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
) {
auto [v_splats, v_viewmat, vr_splats, h_splats] =
_rasterize_to_pixels_eval3d_bwd_tensor<SplatPrimitive, output_distortion, false>
(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height, tile_offsets, flatten_ids,
render_Ts, last_ids, render_outputs, render2_outputs, loss_map,
v_render_outputs, v_render_alphas, v_distortion_outputs,
need_viewmat_grad
);
return std::make_tuple(v_splats, v_viewmat);
}
);
std::tuple<
SphericalVoronoi3DGUT_Default::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_3dgut_sv_bwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render_outputs,
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
);
std::tuple<
VoxelPrimitive::Screen::TensorTuple,
std::optional<at::Tensor> // v_viewmats
> rasterize_to_pixels_voxel_eval3d_bwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// forward outputs
const at::Tensor render_Ts, // [..., image_height, image_width, 1]
const at::Tensor last_ids, // [..., image_height, image_width]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render_outputs,
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> render2_outputs,
std::optional<at::Tensor> loss_map, // [..., image_height, image_width, 1]
// gradients of outputs
VoxelPrimitive::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename VoxelPrimitive::RenderOutput::TensorTuple> v_distortion_outputs,
bool need_viewmat_grad
);
@@ -0,0 +1,463 @@
#pragma once
// Modified from https://github.com/nerfstudio-project/gsplat/blob/main/gsplat/cuda/csrc/RasterizeToPixels3DGSBwd.cu
#include <cuda_runtime.h>
#include <cstdint>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
#ifdef __CUDACC__
#include "generated/slang.cuh"
namespace SlangProjectionUtils {
#include "generated/set_namespace.cuh"
#include "generated/projection_utils.cuh"
}
#endif
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#ifndef NO_TORCH
#define NO_TORCH
#endif
#include "types.cuh"
#include "common.cuh"
// constexpr uint SPLAT_BATCH_SIZE_NO_DISTORTION = 128;
// ^ TODO: behavior similar to https://github.com/nerfstudio-project/gsplat/issues/872
// even if this is a different backward implementation?
constexpr uint SPLAT_BATCH_SIZE_NO_DISTORTION = WARP_SIZE;
// ^ this one is good
constexpr uint SPLAT_BATCH_SIZE_WITH_DISTORTION = WARP_SIZE;
template <
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
bool output_distortion,
bool output_viewmat_grad,
bool output_hessian_diagonal
>
__global__ void rasterize_to_pixels_eval3d_bwd_kernel(
const uint32_t I,
const uint32_t n_isects,
// fwd inputs
typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float *__restrict__ backgrounds, // [..., CDIM] or [nnz, CDIM]
const bool *__restrict__ masks, // [..., tile_height, tile_width]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [..., tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
// fwd outputs
const float *__restrict__ render_Ts, // [..., image_height, image_width, 1]
const int32_t *__restrict__ last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_output_buffer,
typename SplatPrimitive::RenderOutput::Buffer render2_output_buffer,
const float *__restrict__ loss_map_buffer, // [..., image_height, image_width, 1]
// grad outputs
typename SplatPrimitive::RenderOutput::Buffer v_render_output_buffer,
const float *__restrict__ v_render_alphas, // [..., image_height, image_width, 1]
typename SplatPrimitive::RenderOutput::Buffer v_distortions_output_buffer,
// grad inputs
typename SplatPrimitive::Screen::Buffer v_splat_buffer,
typename SplatPrimitive::Screen::Buffer vr_splat_buffer,
typename SplatPrimitive::Screen::Buffer h_splat_buffer,
float *__restrict__ v_viewmats // [B, C, 4, 4]
) {
auto block = cg::this_thread_block();
cg::thread_block_tile<WARP_SIZE> warp = cg::tiled_partition<WARP_SIZE>(block);
uint32_t image_id = block.group_index().x;
uint32_t tile_id = block.group_index().y * tile_width + block.group_index().z;
uint32_t thread_id = block.thread_rank();
tile_offsets += image_id * tile_height * tile_width;
render_Ts += image_id * image_height * image_width;
last_ids += image_id * image_height * image_width;
v_render_alphas += image_id * image_height * image_width;
if (backgrounds != nullptr) {
backgrounds += image_id;
}
if (masks != nullptr) {
masks += image_id * tile_height * tile_width;
}
// when the mask is provided, do nothing and return if
// this tile is labeled as False
if (masks != nullptr && !masks[tile_id]) {
return;
}
// Load camera
viewmats += image_id * 16; // world to camera
float4 intrin = intrins[image_id];
float3x3 R = { // row major
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
CameraDistortionCoeffs dist_coeffs = dist_coeffs_buffer.load(image_id);
constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE;
// load pixels
__shared__ float4 shared_ray_d_pix_bin_final[BLOCK_SIZE];
__shared__ float2 pix_Ts_with_grad[BLOCK_SIZE];
__shared__ typename SplatPrimitive::RenderOutput v_pix_colors[BLOCK_SIZE];
// __shared__ float pix_background[CDIM]; // TODO
__shared__ typename SplatPrimitive::RenderOutput pix_colors[output_distortion ? BLOCK_SIZE : 1];
__shared__ typename SplatPrimitive::RenderOutput pix2_colors[output_distortion ? BLOCK_SIZE : 1];
__shared__ typename SplatPrimitive::RenderOutput v_distortion_out[output_distortion ? BLOCK_SIZE : 1];
__shared__ float hess_weight_map[output_hessian_diagonal ? BLOCK_SIZE : 1];
float3 ray_o = SlangProjectionUtils::transform_ray_o(R, t);
float3 total_v_ray_o = make_float3(0.0f, 0.0f, 0.0f);
__shared__ float3 shared_v_ray_d[output_viewmat_grad ? BLOCK_SIZE : 1];
constexpr uint SPLAT_BATCH_SIZE_CONST = output_distortion ?
SPLAT_BATCH_SIZE_WITH_DISTORTION : SPLAT_BATCH_SIZE_NO_DISTORTION;
#pragma unroll
for (uint pix_id0 = 0; pix_id0 < BLOCK_SIZE; pix_id0 += SPLAT_BATCH_SIZE_CONST) {
static_assert(BLOCK_SIZE % SPLAT_BATCH_SIZE_CONST == 0);
uint pix_id_local = pix_id0 + thread_id;
int pix_x = block.group_index().z * TILE_SIZE + pix_id_local % TILE_SIZE;
int pix_y = block.group_index().y * TILE_SIZE + pix_id_local / TILE_SIZE;
uint pix_id_global = pix_y * image_width + pix_x;
uint pix_id_image_global = image_id * image_height * image_width + pix_id_global;
bool inside = (pix_x < image_width && pix_y < image_height);
int32_t bin_final = (inside ? last_ids[pix_id_global] : 0);
pix_Ts_with_grad[pix_id_local] = {
(inside ? render_Ts[pix_id_global] : 0.0f),
(inside ? -v_render_alphas[pix_id_global] : 0.0f)
};
v_pix_colors[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(v_render_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
const float px = (float)pix_x + 0.5f;
const float py = (float)pix_y + 0.5f;
float3 raydir;
inside &= SlangProjectionUtils::generate_ray(
{(px-cx)/fx, (py-cy)/fy},
camera_model == gsplat::CameraModelType::FISHEYE, dist_coeffs,
&raydir
);
float3 ray_d = SlangProjectionUtils::transform_ray_d(R, raydir); // mul(raydir, R);
shared_ray_d_pix_bin_final[pix_id_local] =
{ray_d.x, ray_d.y, ray_d.z, __int_as_float(bin_final)};
if (output_distortion) {
pix_colors[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(render_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
pix2_colors[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(render2_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
v_distortion_out[pix_id_local] = (inside ?
SplatPrimitive::RenderOutput::load(v_distortions_output_buffer, pix_id_image_global)
: SplatPrimitive::RenderOutput::zero());
}
if (output_viewmat_grad) {
shared_v_ray_d[pix_id_local] = make_float3(0.0f, 0.0f, 0.0f);
}
if (output_hessian_diagonal) {
// https://www.desmos.com/calculator/ld9wg7cuxz
hess_weight_map[pix_id_local] = (loss_map_buffer != nullptr && inside) ?
0.5f / fmaxf(loss_map_buffer[pix_id_image_global], 1e-6f) : 0.0f;
}
}
block.sync();
// threads fist load splats, then swept through pixels
// do this in batches
int32_t range_start = tile_offsets[tile_id];
int32_t range_end =
(image_id == I - 1) && (tile_id == tile_width * tile_height - 1)
? n_isects
: tile_offsets[tile_id + 1];
uint SPLAT_BATCH_SIZE = SPLAT_BATCH_SIZE_CONST;
if (SPLAT_BATCH_SIZE_CONST > WARP_SIZE) {
SPLAT_BATCH_SIZE = (uint)sqrtf((float)(range_end - range_start) * (float)BLOCK_SIZE);
// SPLAT_BATCH_SIZE = min(SPLAT_BATCH_SIZE_CONST, (SPLAT_BATCH_SIZE + WARP_SIZE) & ~(WARP_SIZE-1));
SPLAT_BATCH_SIZE = min(SPLAT_BATCH_SIZE_CONST, max(SPLAT_BATCH_SIZE, 1u));
}
const uint32_t num_splat_batches =
_CEIL_DIV(range_end - range_start, SPLAT_BATCH_SIZE);
// if (warp.thread_rank() == 0)
// printf("range_start=%d range_end=%d num_splat_batches=%u\n", range_start, range_end, num_splat_batches);
for (uint32_t splat_b = 0; splat_b < num_splat_batches; ++splat_b) {
const int32_t splat_batch_end = range_end - 1 - SPLAT_BATCH_SIZE * splat_b;
const int32_t splat_batch_size = min(SPLAT_BATCH_SIZE, splat_batch_end + 1 - range_start);
const int32_t splat_idx = splat_batch_end - thread_id;
// load splats
typename SplatPrimitive::Screen splat;
uint32_t splat_gid;
if (splat_idx >= range_start) {
splat_gid = flatten_ids[splat_idx]; // flatten index in [I * N] or [nnz]
splat = SplatPrimitive::Screen::loadWithPrecompute(splat_buffer, splat_gid);
}
// accumulate gradient
typename SplatPrimitive::Screen v_splat = SplatPrimitive::Screen::zero();
typename SplatPrimitive::Screen vr_splat = SplatPrimitive::Screen::zero();
typename SplatPrimitive::Screen h_splat = SplatPrimitive::Screen::zero();
// thread 0 takes last splat, 1 takes second last, etc.
// at t=0, thread 0 (splat -1) undo pixel 0
// at t=1, thread 0 (splat -1) undo pixel 1, thread 1 (splat -2) undo pixel 0
// ......
// process gaussians in the current batch for this pixel
// 0 index is the furthest back gaussian in the batch
for (int t = 0; t < splat_batch_size + BLOCK_SIZE - 1; ++t, __syncwarp()) {
int pix_id = t - thread_id;
if (pix_id < 0 || pix_id >= BLOCK_SIZE || splat_idx < range_start)
continue;
float4 ray_d_pix_bin_final = shared_ray_d_pix_bin_final[pix_id];
if (splat_idx > __float_as_int(ray_d_pix_bin_final.w))
continue;
// evaluate alpha and early skip
float3 ray_d = {ray_d_pix_bin_final.x, ray_d_pix_bin_final.y, ray_d_pix_bin_final.z};
float alpha = splat.evaluate_alpha(ray_o, ray_d);
if (alpha <= ALPHA_THRESHOLD || dot(ray_d, ray_d) == 0.0f)
continue;
// printf("t=%d, thread %u, splat %d (%u), pix_id %d, pix %d %d\n", t, thread_id, splat_idx-range_start, splat_gid, pix_id, pix_global_x, pix_global_y);
// forward:
// \left(c_{1},T_{1}\right)=\left(c_{0}+\alpha_{i}T_{0}c_{i},\ T_{0}\left(1-\alpha_{i}\right)\right)
float T1 = pix_Ts_with_grad[pix_id].x;
float v_T1 = pix_Ts_with_grad[pix_id].y;
// undo pixel:
// T_{0}=\frac{T_{1}}{1-\alpha_{i}}
float ra = 1.0f / (1.0f - alpha);
float T0 = T1 * ra;
typename SplatPrimitive::RenderOutput color = splat.evaluate_color(ray_o, ray_d);
typename SplatPrimitive::RenderOutput v_c = v_pix_colors[pix_id];
// gradient to alpha:
// \frac{dL}{d\alpha_{i}}
// = \frac{dL}{dc_{1}}\frac{dc_{1}}{d\alpha_{i}}+\frac{dL}{dT_{1}}\frac{dT_{1}}{d\alpha_{i}}
// = T_{0}\frac{dL}{dc_{1}}c_{i}-\frac{dL}{dT_{1}}T_{0}
float v_alpha = T0 * color.dot(v_c) -v_T1 * T0;
// gradient to color:
// \frac{dL}{dc_{i}}
// = \frac{dL}{dc_{1}}\frac{dc_{1}}{dc_{i}}
// = \alpha_{i}T_{0}\frac{dL}{dc_{1}}
typename SplatPrimitive::RenderOutput v_color = v_c * (alpha * T0);
// update pixel gradient:
// \frac{dL}{dT_{0}}
// = \frac{dL}{dc_{1}}\frac{dc_{1}}{dT_{0}}+\frac{dL}{dT_{1}}\frac{dT_{1}}{dT_{0}}
// = \alpha_{i}\frac{dL}{dc_{1}}c_{i}+\frac{dL}{dT_{1}}\left(1-\alpha_{i}\right)
float v_T0 = alpha * color.dot(v_c) + v_T1 * (1.0f - alpha);
// distortion
if (output_distortion) {
// \left(d_{1},s_{1}\right)=\left(d_{0}+\alpha_{i}T_{0}\left(c_{i}^{2}\left(1-T_{0}\right)-2c_{i}c_{0}+s_{0}\right),\ s_{0}+\alpha_{i}T_{0}c_{i}^{2}\right)
// \frac{dL}{ds}=0
typename SplatPrimitive::RenderOutput v_dist = v_distortion_out[pix_id];
typename SplatPrimitive::RenderOutput c0 =
pix_colors[pix_id] + color * -alpha * T0;
typename SplatPrimitive::RenderOutput s0 =
pix2_colors[pix_id] + color * color * -alpha * T0;
// \frac{dL}{d\alpha_{i}}=\frac{dL}{dd_{1}}\frac{dd_{1}}{d\alpha_{i}}=T_{0}\left(c_{i}^{2}\left(1-T_{0}\right)-2c_{i}c_{0}+s_{0}\right)\frac{dL}{dd_{1}}
v_alpha += T0 * (
color * color * (1.0f-T0) +
color * c0 * -2.0f + s0
).dot(v_dist);
// \frac{dL}{dc_{i}}=\frac{dL}{dd_{1}}\frac{dd_{1}}{dc_{i}}=2\alpha_{i}T_{0}\left(c_{i}\left(1-T_{0}\right)-c_{0}\right)\frac{dL}{dd_{1}}
v_color += (
color * (1.0f-T0) +
c0 * -1.0f
) * v_dist * (2.0f * alpha * T0);
// \alpha_{i}\left(c_{i}^{2}\left(1-2T_{0}\right)-2c_{i}c_{0}+s_{0}\right)\frac{dL}{dd_{1}}
v_T0 += alpha * (
color * color * (1.0f-2.0f*T0) +
color * c0 * -2.0f + s0
).dot(v_dist);
// undo pixel state
pix_colors[pix_id] = c0;
pix2_colors[pix_id] = s0;
}
// backward diff splat
float3 v_ray_o_alpha, v_ray_d_alpha;
float3 v_ray_o_color, v_ray_d_color;
if (output_hessian_diagonal) {
typename SplatPrimitive::Screen v_splat_temp = SplatPrimitive::Screen::zero();
v_splat_temp.addGradient(splat.evaluate_alpha_vjp(ray_o, ray_d, v_alpha, v_ray_o_alpha, v_ray_d_alpha));
v_splat_temp.addGradient(splat.evaluate_color_vjp(ray_o, ray_d, v_color, v_ray_o_color, v_ray_d_color));
v_splat.addGradient(v_splat_temp);
// https://www.desmos.com/calculator/ld9wg7cuxz
float weight = 1.0f / sqrtf(fmaxf(v_c.dot(v_c) / 3.0f, 1e-30f));
vr_splat.addGradient(v_splat_temp, weight);
weight = hess_weight_map[pix_id] * weight * weight;
splat.precomputeBackward(v_splat_temp);
h_splat.addGaussNewtonHessianDiagonal(v_splat_temp, weight);
} else {
v_splat.addGradient(splat.evaluate_alpha_vjp(ray_o, ray_d, v_alpha, v_ray_o_alpha, v_ray_d_alpha));
v_splat.addGradient(splat.evaluate_color_vjp(ray_o, ray_d, v_color, v_ray_o_color, v_ray_d_color));
}
if (output_viewmat_grad) {
total_v_ray_o += v_ray_o_alpha + v_ray_o_color;
shared_v_ray_d[pix_id] += v_ray_d_alpha + v_ray_d_color;
}
// update pixel states
pix_Ts_with_grad[pix_id] = { T0, v_T0 };
// v_pix_colors remains the same
}
// accumulate gradient
if (splat_idx >= range_start) {
splat.precomputeBackward(v_splat);
v_splat.atomicAddToBuffer(v_splat_buffer, splat_gid);
if (output_hessian_diagonal) {
splat.precomputeBackward(vr_splat);
vr_splat.atomicAddToBuffer(vr_splat_buffer, splat_gid);
h_splat.atomicAddToBuffer(h_splat_buffer, splat_gid);
}
}
}
if (output_viewmat_grad) {
// accumulate to viewmat gradient
float3x3 v_R;
float3 v_t;
// gradient from ray_o (will fill v_R and v_t)
SlangProjectionUtils::transform_ray_o_vjp(R, t, total_v_ray_o, &v_R, &v_t);
// gradient from ray_d
#pragma unroll
for (uint pix_id0 = 0; pix_id0 < BLOCK_SIZE; pix_id0 += SPLAT_BATCH_SIZE_CONST) {
uint pix_id_local = pix_id0 + thread_id;
float4 ray_d_pix_bin_final = shared_ray_d_pix_bin_final[pix_id_local];
float3 raydir = SlangProjectionUtils::undo_transform_ray_d(R,
make_float3(
ray_d_pix_bin_final.x,
ray_d_pix_bin_final.y,
ray_d_pix_bin_final.z
)
);
float3 v_ray_d = shared_v_ray_d[pix_id_local];
float3x3 v_R_delta;
float3 temp;
SlangProjectionUtils::transform_ray_d_vjp(R, raydir, v_ray_d, &v_R_delta, &temp);
v_R = v_R + v_R_delta;
}
// atomic add to global viewmat gradient
if (v_viewmats != nullptr) {
float *v_viewmat = v_viewmats + image_id * 16;
float temp;
#define _ATOMIC_ADD(ptr, offset, value) do { \
temp = isfinite(value) ? value : 0.0f; \
warpSum(temp, warp); \
if (warp.thread_rank() == 0 && temp != 0.0f) \
atomicAdd((ptr) + (offset), (temp)); \
} while(0)
_ATOMIC_ADD(v_viewmat, 0, v_R[0].x);
_ATOMIC_ADD(v_viewmat, 1, v_R[0].y);
_ATOMIC_ADD(v_viewmat, 2, v_R[0].z);
_ATOMIC_ADD(v_viewmat, 3, v_t.x);
_ATOMIC_ADD(v_viewmat, 4, v_R[1].x);
_ATOMIC_ADD(v_viewmat, 5, v_R[1].y);
_ATOMIC_ADD(v_viewmat, 6, v_R[1].z);
_ATOMIC_ADD(v_viewmat, 7, v_t.y);
_ATOMIC_ADD(v_viewmat, 8, v_R[2].x);
_ATOMIC_ADD(v_viewmat, 9, v_R[2].y);
_ATOMIC_ADD(v_viewmat, 10, v_R[2].z);
_ATOMIC_ADD(v_viewmat, 11, v_t.z);
#undef _ATOMIC_ADD
}
}
}
template <
typename SplatPrimitive,
gsplat::CameraModelType camera_model,
bool output_distortion,
bool output_viewmat_grad,
bool output_hessian_diagonal
>
void rasterize_to_pixels_eval3d_bwd_kernel_wrapper(
cudaStream_t stream,
const uint32_t I,
const uint32_t n_isects,
// fwd inputs
typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float *__restrict__ backgrounds, // [..., CDIM] or [nnz, CDIM]
const bool *__restrict__ masks, // [..., tile_height, tile_width]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [..., tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
// fwd outputs
const float *__restrict__ render_Ts, // [..., image_height, image_width, 1]
const int32_t *__restrict__ last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_output_buffer,
typename SplatPrimitive::RenderOutput::Buffer render2_output_buffer,
const float *__restrict__ loss_map_buffer, // [..., image_height, image_width, 1]
// grad outputs
typename SplatPrimitive::RenderOutput::Buffer v_render_output_buffer,
const float *__restrict__ v_render_alphas, // [..., image_height, image_width, 1]
typename SplatPrimitive::RenderOutput::Buffer v_distortions_output_buffer,
// grad inputs
typename SplatPrimitive::Screen::Buffer v_splat_buffer,
typename SplatPrimitive::Screen::Buffer vr_splat_buffer,
typename SplatPrimitive::Screen::Buffer h_splat_buffer,
float *__restrict__ v_viewmats // [B, C, 4, 4]
) {
dim3 threads = {output_distortion ?
SPLAT_BATCH_SIZE_WITH_DISTORTION : SPLAT_BATCH_SIZE_NO_DISTORTION,
1, 1};
dim3 grid = {I, tile_height, tile_width};
rasterize_to_pixels_eval3d_bwd_kernel<
SplatPrimitive, camera_model, output_distortion, output_viewmat_grad, output_hessian_diagonal
><<<grid, threads, 0, stream>>>(
I, n_isects,
splat_buffer,
viewmats, intrins, dist_coeffs_buffer, backgrounds, masks,
image_width, image_height, tile_width, tile_height,
tile_offsets, flatten_ids,
render_Ts, last_ids,
render_output_buffer, render2_output_buffer, loss_map_buffer,
v_render_output_buffer, v_render_alphas,
v_distortions_output_buffer, v_splat_buffer, vr_splat_buffer, h_splat_buffer,
v_viewmats
);
}
@@ -0,0 +1,313 @@
#include "RasterizationEval3DFwd.cuh"
#include <gsplat/Utils.cuh>
#include "types.cuh"
#include "common.cuh"
#include <c10/cuda/CUDAStream.h>
template <typename SplatPrimitive, gsplat::CameraModelType camera_model, bool output_distortion, bool output_max_blending>
void rasterize_to_pixels_eval3d_fwd_kernel_wrapper(
cudaStream_t stream,
const uint32_t I,
const uint32_t N,
const uint32_t n_isects,
const bool packed,
const typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float3 *__restrict__ backgrounds, // [I, 3]
const bool *__restrict__ max_blending_masks, // [B, C, image_width, image_height]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [I, tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
typename SplatPrimitive::RenderOutput::Buffer render_colors, // [I, image_height, image_width, ...]
float *__restrict__ render_Ts, // [I, image_height, image_width, 1]
int32_t *__restrict__ last_ids, // [I, image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...]
typename SplatPrimitive::RenderOutput::Buffer render_distortions, // [I, image_height, image_width, ...]
float* __restrict__ out_max_blending
);
template <typename SplatPrimitive, bool output_distortion, bool output_max_blending>
inline void launch_rasterize_to_pixels_eval3d_fwd_kernel(
// Gaussian parameters
typename SplatPrimitive::Screen::Tensor splats,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> max_blending_masks, // [..., C, image_width, image_height]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// outputs
typename SplatPrimitive::RenderOutput::Tensor renders,
at::Tensor transmittances, // [..., image_height, image_width]
at::Tensor last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Tensor *renders2,
typename SplatPrimitive::RenderOutput::Tensor *distortions,
std::optional<at::Tensor>& out_max_blending
) {
bool packed = splats.isPacked();
uint32_t N = packed ? 0 : splats.size(); // number of gaussians
uint32_t I = transmittances.numel() / (image_height * image_width); // number of images
uint32_t tile_height = tile_offsets.size(-2);
uint32_t tile_width = tile_offsets.size(-1);
uint32_t n_isects = flatten_ids.size(0);
#define _LAUNCH_ARGS ( \
(cudaStream_t)at::cuda::getCurrentCUDAStream(), I, N, n_isects, packed, \
splats.buffer(), \
viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
backgrounds.has_value() ? (float3*)backgrounds.value().data_ptr<float>() : nullptr, \
(output_max_blending && max_blending_masks.has_value()) ? max_blending_masks.value().data_ptr<bool>() : nullptr, \
image_width, image_height, tile_width, tile_height, \
tile_offsets.data_ptr<int32_t>(), flatten_ids.data_ptr<int32_t>(), \
renders, transmittances.data_ptr<float>(), last_ids.data_ptr<int32_t>(), \
output_distortion ? renders2->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
output_distortion ? distortions->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
(output_max_blending && out_max_blending.has_value()) ? out_max_blending.value().data_ptr<float>() : nullptr \
)
if (camera_model == gsplat::CameraModelType::PINHOLE)
rasterize_to_pixels_eval3d_fwd_kernel_wrapper<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion, output_max_blending> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
rasterize_to_pixels_eval3d_fwd_kernel_wrapper<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion, output_max_blending> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
}
template <typename SplatPrimitive, bool output_distortion, bool output_max_blending>
inline std::tuple<
typename SplatPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_eval3d_fwd_tensor(
// Gaussian parameters
typename SplatPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> max_blending_masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids // [n_isects]
) {
DEVICE_GUARD(tile_offsets);
CHECK_INPUT(tile_offsets);
CHECK_INPUT(flatten_ids);
CHECK_INPUT(viewmats);
CHECK_INPUT(intrins);
if (backgrounds.has_value())
CHECK_INPUT(backgrounds.value());
if (output_max_blending && max_blending_masks.has_value())
CHECK_INPUT(max_blending_masks.value());
typename SplatPrimitive::Screen::Tensor splats(splats_tuple);
auto opt = splats.options();
at::DimVector image_dims(tile_offsets.sizes().slice(0, tile_offsets.dim() - 2));
at::DimVector renders_dims(image_dims);
renders_dims.append({image_height, image_width});
typename SplatPrimitive::RenderOutput::Tensor renders =
SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt);
std::optional<typename SplatPrimitive::RenderOutput::Tensor> renders2 = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> distortions = std::nullopt;
if (output_distortion) {
renders2 = SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt);
distortions = SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt);
}
at::DimVector transmittance_dims(image_dims);
transmittance_dims.append({image_height, image_width, 1});
at::Tensor transmittances = at::empty(transmittance_dims, opt);
at::DimVector last_ids_dims(image_dims);
last_ids_dims.append({image_height, image_width});
at::Tensor last_ids = at::empty(last_ids_dims, opt.dtype(at::kInt));
std::optional<at::Tensor> out_max_blending;
if (output_max_blending) {
out_max_blending = at::empty({splats.size()}, opt);
set_zero<float>(out_max_blending.value());
}
launch_rasterize_to_pixels_eval3d_fwd_kernel<SplatPrimitive, output_distortion, output_max_blending>(
splats,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, max_blending_masks,
image_width, image_height, tile_offsets, flatten_ids,
renders, transmittances, last_ids,
output_distortion ? &renders2.value() : nullptr,
output_distortion ? &distortions.value() : nullptr,
out_max_blending
);
if (output_distortion)
return std::make_tuple(renders.tuple(), transmittances, last_ids,
renders2.value().tuple(), distortions.value().tuple(), out_max_blending);
return std::make_tuple(renders.tuple(), transmittances, last_ids,
std::nullopt, std::nullopt, out_max_blending);
}
// ================
// Vanilla3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
Vanilla3DGUT::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_fwd(
// Gaussian parameters
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
) {
if (output_distortion)
return rasterize_to_pixels_eval3d_fwd_tensor<Vanilla3DGUT, true, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
return rasterize_to_pixels_eval3d_fwd_tensor<Vanilla3DGUT, false, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
// ================
// SphericalVoronoi3DGUT
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_sv_fwd(
// Gaussian parameters
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
) {
if (output_distortion)
return rasterize_to_pixels_eval3d_fwd_tensor<SphericalVoronoi3DGUT_Default, true, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
return rasterize_to_pixels_eval3d_fwd_tensor<SphericalVoronoi3DGUT_Default, false, false>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
// ================
// VoxelPrimitive
// ================
/*[AutoHeaderGeneratorExport]*/
std::tuple<
VoxelPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_voxel_eval3d_fwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids // [n_isects]
) {
return rasterize_to_pixels_eval3d_fwd_tensor<VoxelPrimitive, true, true>(
splats_tuple,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
@@ -1,356 +1,97 @@
#pragma once
// Modified from https://github.com/nerfstudio-project/gsplat/blob/main/gsplat/cuda/csrc/RasterizeToPixels3DGSFwd.cu
#include <cuda.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <ATen/Tensor.h>
#ifdef __CUDACC__
#include "generated/slang.cuh"
namespace SlangProjectionUtils {
#include "generated/set_namespace.cuh"
#include "generated/projection_utils.cuh"
}
#endif
// #include "Primitive3DGS.cuh"
#include "Primitive3DGUT.cuh"
#include "Primitive3DGUT_SV.cuh"
// #include "PrimitiveOpaqueTriangle.cuh"
#include "PrimitiveVoxel.cuh"
#include "types.cuh"
#include "common.cuh"
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
template <typename SplatPrimitive, gsplat::CameraModelType camera_model, bool output_distortion, bool output_max_blending>
__global__ void rasterize_to_pixels_eval3d_fwd_kernel(
const uint32_t I,
const uint32_t N,
const uint32_t n_isects,
const bool packed,
const typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float3 *__restrict__ backgrounds, // [I, 3]
const bool *__restrict__ max_blending_masks, // [B, C, image_width, image_height]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [I, tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
typename SplatPrimitive::RenderOutput::Buffer render_colors, // [I, image_height, image_width, ...]
float *__restrict__ render_Ts, // [I, image_height, image_width, 1]
int32_t *__restrict__ last_ids, // [I, image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...]
typename SplatPrimitive::RenderOutput::Buffer render_distortions, // [I, image_height, image_width, ...]
float* __restrict__ out_max_blending
) {
// each thread draws one pixel, but also timeshares caching gaussians in a
// shared tile
auto block = cg::this_thread_block();
int32_t image_id = block.group_index().x;
int32_t tile_id =
block.group_index().y * tile_width + block.group_index().z;
uint32_t i = block.group_index().y * TILE_SIZE + block.thread_index().y;
uint32_t j = block.group_index().z * TILE_SIZE + block.thread_index().x;
tile_offsets += image_id * tile_height * tile_width;
render_Ts += image_id * image_height * image_width;
last_ids += image_id * image_height * image_width;
if (backgrounds != nullptr) {
backgrounds += image_id;
}
if (output_max_blending && max_blending_masks != nullptr)
max_blending_masks += image_id * image_height * image_width;
float px = (float)j + 0.5f;
float py = (float)i + 0.5f;
int32_t pix_id = i * image_width + j;
// Load camera
viewmats += image_id * 16;
float4 intrin = intrins[image_id];
float3x3 R = {
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
CameraDistortionCoeffs dist_coeffs = dist_coeffs_buffer.load(image_id);
bool inside = (i < image_height && j < image_width);
float3 raydir;
inside &= SlangProjectionUtils::generate_ray(
{(px-cx)/fx, (py-cy)/fy},
camera_model == gsplat::CameraModelType::FISHEYE, dist_coeffs,
&raydir
);
float3 ray_o = SlangProjectionUtils::transform_ray_o(R, t);
float3 ray_d = SlangProjectionUtils::transform_ray_d(R, raydir);
bool done = !inside;
// have all threads in tile process the same gaussians in batches
// first collect gaussians between range.x and range.y in batches
// which gaussians to look through in this tile
int32_t range_start = tile_offsets[tile_id];
int32_t range_end =
(image_id == I - 1) && (tile_id == tile_width * tile_height - 1)
? n_isects
: tile_offsets[tile_id + 1];
const uint32_t block_size = block.size();
uint32_t num_batches =
(range_end - range_start + block_size - 1) / block_size;
constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE;
__shared__ typename SplatPrimitive::Screen splat_batch[BLOCK_SIZE];
__shared__ uint32_t splat_idx_batch[output_max_blending ? BLOCK_SIZE : 1];
// current visibility left to render
// transmittance is gonna be used in the backward pass which requires a high
// numerical precision so we use double for it. However double make bwd 1.5x
// slower so we stick with float for now.
float T = 1.0f;
// index of most recent gaussian to write to this thread's pixel
uint32_t cur_idx = 0;
// collect and process batches of gaussians
// each thread loads one gaussian at a time before rasterizing its
// designated pixel
uint32_t tr = block.thread_rank();
bool max_blending_mask = (output_max_blending && max_blending_masks) ?
max_blending_masks[pix_id] : true;
typename SplatPrimitive::RenderOutput pix_out = SplatPrimitive::RenderOutput::zero();
typename SplatPrimitive::RenderOutput pix2_out = SplatPrimitive::RenderOutput::zero();
typename SplatPrimitive::RenderOutput distortion_out = SplatPrimitive::RenderOutput::zero();
for (uint32_t b = 0; b < num_batches; ++b) {
// resync all threads before beginning next batch
// end early if entire tile is done
if (__syncthreads_count(done) >= block_size) {
break;
}
// each thread fetch 1 gaussian from front to back
// index of gaussian to load
uint32_t batch_start = range_start + block_size * b;
uint32_t idx = batch_start + tr;
if (idx < range_end) {
int32_t g = flatten_ids[idx]; // flatten index in [I * N] or [nnz]
splat_batch[tr] = SplatPrimitive::Screen::loadWithPrecompute(splat_buffer, g);
if (output_max_blending)
splat_idx_batch[tr] = g;
}
// wait for other threads to collect the gaussians in batch
block.sync();
// process gaussians in the current batch for this pixel
uint32_t batch_size = min(block_size, range_end - batch_start);
for (uint32_t t = 0; (t < batch_size) && !done; ++t) {
typename SplatPrimitive::Screen splat = splat_batch[t];
float alpha = splat.evaluate_alpha(ray_o, ray_d);
if (alpha < ALPHA_THRESHOLD) {
continue;
}
const float next_T = T * (1.0f - alpha);
if (next_T <= 1e-4f) { // this pixel is done: exclusive
done = true;
break;
}
const float vis = alpha * T;
const typename SplatPrimitive::RenderOutput color = splat.evaluate_color(ray_o, ray_d);
if (output_distortion) {
distortion_out += (
color * color * (1.0f - T)
+ color * pix_out * -2.0f
+ pix2_out
) * vis;
pix2_out += color * color * vis;
}
pix_out += color * vis;
cur_idx = batch_start + t;
T = next_T;
if (output_max_blending && out_max_blending != nullptr && max_blending_mask) {
uint32_t splat_idx = splat_idx_batch[t];
atomicMax(out_max_blending + (N == 0 ? splat_idx : splat_idx % N), vis);
}
}
}
if (i < image_height && j < image_width) {
render_Ts[pix_id] = T;
int pix_id_global = image_id * image_height * image_width + pix_id;
// TODO: blend background
pix_out.saveParamsToBuffer(render_colors, pix_id_global);
// index in bin of last gaussian in this pixel
last_ids[pix_id] = static_cast<int32_t>(cur_idx);
// distortion
if (output_distortion) {
pix2_out.saveParamsToBuffer(render_colors2, pix_id_global);
distortion_out.saveParamsToBuffer(render_distortions, pix_id_global);
}
}
}
template <typename SplatPrimitive, bool output_distortion, bool output_max_blending>
inline void launch_rasterize_to_pixels_eval3d_fwd_kernel(
std::tuple<
Vanilla3DGUT::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<Vanilla3DGUT::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_3dgut_fwd(
// Gaussian parameters
typename SplatPrimitive::Screen::Tensor splats,
Vanilla3DGUT::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> max_blending_masks, // [..., C, image_width, image_height]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
// outputs
typename SplatPrimitive::RenderOutput::Tensor renders,
at::Tensor transmittances, // [..., image_height, image_width]
at::Tensor last_ids, // [..., image_height, image_width]
typename SplatPrimitive::RenderOutput::Tensor *renders2,
typename SplatPrimitive::RenderOutput::Tensor *distortions,
std::optional<at::Tensor>& out_max_blending
) {
bool packed = splats.isPacked();
uint32_t N = packed ? 0 : splats.size(); // number of gaussians
uint32_t I = transmittances.numel() / (image_height * image_width); // number of images
uint32_t tile_height = tile_offsets.size(-2);
uint32_t tile_width = tile_offsets.size(-1);
uint32_t n_isects = flatten_ids.size(0);
// Each block covers a tile on the image. In total there are
// I * tile_height * tile_width blocks.
dim3 threads = {TILE_SIZE, TILE_SIZE, 1};
dim3 grid = {I, tile_height, tile_width};
#define _LAUNCH_ARGS <<<grid, threads, 0, at::cuda::getCurrentCUDAStream()>>>( \
I, N, n_isects, packed, \
splats.buffer(), \
viewmats.data_ptr<float>(), (float4*)intrins.data_ptr<float>(), dist_coeffs, \
backgrounds.has_value() ? (float3*)backgrounds.value().data_ptr<float>() : nullptr, \
(output_max_blending && max_blending_masks.has_value()) ? max_blending_masks.value().data_ptr<bool>() : nullptr, \
image_width, image_height, tile_width, tile_height, \
tile_offsets.data_ptr<int32_t>(), flatten_ids.data_ptr<int32_t>(), \
renders, transmittances.data_ptr<float>(), last_ids.data_ptr<int32_t>(), \
output_distortion ? renders2->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
output_distortion ? distortions->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \
(output_max_blending && out_max_blending.has_value()) ? out_max_blending.value().data_ptr<float>() : nullptr \
)
if (camera_model == gsplat::CameraModelType::PINHOLE)
rasterize_to_pixels_eval3d_fwd_kernel<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion, output_max_blending> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
rasterize_to_pixels_eval3d_fwd_kernel<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion, output_max_blending> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
}
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
);
template <typename SplatPrimitive, bool output_distortion, bool output_max_blending>
inline std::tuple<
typename SplatPrimitive::RenderOutput::TensorTuple,
std::tuple<
SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<SphericalVoronoi3DGUT_Default::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_eval3d_fwd_tensor(
> rasterize_to_pixels_3dgut_sv_fwd(
// Gaussian parameters
typename SplatPrimitive::Screen::TensorTuple splats_tuple,
SphericalVoronoi3DGUT_Default::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> max_blending_masks, // [..., tile_height, tile_width]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids, // [n_isects]
bool output_distortion
);
std::tuple<
VoxelPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<VoxelPrimitive::RenderOutput::TensorTuple>,
std::optional<at::Tensor>
> rasterize_to_pixels_voxel_eval3d_fwd(
// Gaussian parameters
VoxelPrimitive::Screen::TensorTuple splats_tuple,
const at::Tensor viewmats, // [..., C, 4, 4]
const at::Tensor intrins, // [..., C, 4], fx, fy, cx, cy
const gsplat::CameraModelType camera_model,
const CameraDistortionCoeffsTensor dist_coeffs,
const std::optional<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> masks, // [..., tile_height, tile_width]
// image size
const uint32_t image_width,
const uint32_t image_height,
// intersections
const at::Tensor tile_offsets, // [..., tile_height, tile_width]
const at::Tensor flatten_ids // [n_isects]
) {
DEVICE_GUARD(tile_offsets);
CHECK_INPUT(tile_offsets);
CHECK_INPUT(flatten_ids);
CHECK_INPUT(viewmats);
CHECK_INPUT(intrins);
if (backgrounds.has_value())
CHECK_INPUT(backgrounds.value());
if (output_max_blending && max_blending_masks.has_value())
CHECK_INPUT(max_blending_masks.value());
typename SplatPrimitive::Screen::Tensor splats(splats_tuple);
auto opt = splats.options();
at::DimVector image_dims(tile_offsets.sizes().slice(0, tile_offsets.dim() - 2));
at::DimVector renders_dims(image_dims);
renders_dims.append({image_height, image_width});
typename SplatPrimitive::RenderOutput::Tensor renders =
SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt);
std::optional<typename SplatPrimitive::RenderOutput::Tensor> renders2 = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> distortions = std::nullopt;
if (output_distortion) {
renders2 = SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt);
distortions = SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt);
}
at::DimVector transmittance_dims(image_dims);
transmittance_dims.append({image_height, image_width, 1});
at::Tensor transmittances = at::empty(transmittance_dims, opt);
at::DimVector last_ids_dims(image_dims);
last_ids_dims.append({image_height, image_width});
at::Tensor last_ids = at::empty(last_ids_dims, opt.dtype(at::kInt));
std::optional<at::Tensor> out_max_blending;
if (output_max_blending) {
out_max_blending = at::empty({splats.size()}, opt);
set_zero<float>(out_max_blending.value());
}
launch_rasterize_to_pixels_eval3d_fwd_kernel<SplatPrimitive, output_distortion, output_max_blending>(
splats,
viewmats, intrins, camera_model, dist_coeffs,
backgrounds, max_blending_masks,
image_width, image_height, tile_offsets, flatten_ids,
renders, transmittances, last_ids,
output_distortion ? &renders2.value() : nullptr,
output_distortion ? &distortions.value() : nullptr,
out_max_blending
);
if (output_distortion)
return std::make_tuple(renders.tuple(), transmittances, last_ids,
renders2.value().tuple(), distortions.value().tuple(), out_max_blending);
return std::make_tuple(renders.tuple(), transmittances, last_ids,
std::nullopt, std::nullopt, out_max_blending);
}
);
@@ -0,0 +1,251 @@
#pragma once
// Modified from https://github.com/nerfstudio-project/gsplat/blob/main/gsplat/cuda/csrc/RasterizeToPixels3DGSFwd.cu
#include <cuda.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <cooperative_groups.h>
namespace cg = cooperative_groups;
#ifdef __CUDACC__
#include "generated/slang.cuh"
namespace SlangProjectionUtils {
#include "generated/set_namespace.cuh"
#include "generated/projection_utils.cuh"
}
#endif
#include "types.cuh"
#include "common.cuh"
#include <gsplat/Common.h>
#include <gsplat/Utils.cuh>
template <typename SplatPrimitive, gsplat::CameraModelType camera_model, bool output_distortion, bool output_max_blending>
__global__ void rasterize_to_pixels_eval3d_fwd_kernel(
const uint32_t I,
const uint32_t N,
const uint32_t n_isects,
const bool packed,
const typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float3 *__restrict__ backgrounds, // [I, 3]
const bool *__restrict__ max_blending_masks, // [B, C, image_width, image_height]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [I, tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
typename SplatPrimitive::RenderOutput::Buffer render_colors, // [I, image_height, image_width, ...]
float *__restrict__ render_Ts, // [I, image_height, image_width, 1]
int32_t *__restrict__ last_ids, // [I, image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...]
typename SplatPrimitive::RenderOutput::Buffer render_distortions, // [I, image_height, image_width, ...]
float* __restrict__ out_max_blending
) {
// each thread draws one pixel, but also timeshares caching gaussians in a
// shared tile
auto block = cg::this_thread_block();
int32_t image_id = block.group_index().x;
int32_t tile_id =
block.group_index().y * tile_width + block.group_index().z;
uint32_t i = block.group_index().y * TILE_SIZE + block.thread_index().y;
uint32_t j = block.group_index().z * TILE_SIZE + block.thread_index().x;
tile_offsets += image_id * tile_height * tile_width;
render_Ts += image_id * image_height * image_width;
last_ids += image_id * image_height * image_width;
if (backgrounds != nullptr) {
backgrounds += image_id;
}
if (output_max_blending && max_blending_masks != nullptr)
max_blending_masks += image_id * image_height * image_width;
float px = (float)j + 0.5f;
float py = (float)i + 0.5f;
int32_t pix_id = i * image_width + j;
// Load camera
viewmats += image_id * 16;
float4 intrin = intrins[image_id];
float3x3 R = {
viewmats[0], viewmats[1], viewmats[2], // 1st row
viewmats[4], viewmats[5], viewmats[6], // 2nd row
viewmats[8], viewmats[9], viewmats[10], // 3rd row
};
float3 t = { viewmats[3], viewmats[7], viewmats[11] };
float fx = intrin.x, fy = intrin.y, cx = intrin.z, cy = intrin.w;
CameraDistortionCoeffs dist_coeffs = dist_coeffs_buffer.load(image_id);
bool inside = (i < image_height && j < image_width);
float3 raydir;
inside &= SlangProjectionUtils::generate_ray(
{(px-cx)/fx, (py-cy)/fy},
camera_model == gsplat::CameraModelType::FISHEYE, dist_coeffs,
&raydir
);
float3 ray_o = SlangProjectionUtils::transform_ray_o(R, t);
float3 ray_d = SlangProjectionUtils::transform_ray_d(R, raydir);
bool done = !inside;
// have all threads in tile process the same gaussians in batches
// first collect gaussians between range.x and range.y in batches
// which gaussians to look through in this tile
int32_t range_start = tile_offsets[tile_id];
int32_t range_end =
(image_id == I - 1) && (tile_id == tile_width * tile_height - 1)
? n_isects
: tile_offsets[tile_id + 1];
const uint32_t block_size = block.size();
uint32_t num_batches =
(range_end - range_start + block_size - 1) / block_size;
constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE;
__shared__ typename SplatPrimitive::Screen splat_batch[BLOCK_SIZE];
__shared__ uint32_t splat_idx_batch[output_max_blending ? BLOCK_SIZE : 1];
// current visibility left to render
// transmittance is gonna be used in the backward pass which requires a high
// numerical precision so we use double for it. However double make bwd 1.5x
// slower so we stick with float for now.
float T = 1.0f;
// index of most recent gaussian to write to this thread's pixel
uint32_t cur_idx = 0;
// collect and process batches of gaussians
// each thread loads one gaussian at a time before rasterizing its
// designated pixel
uint32_t tr = block.thread_rank();
bool max_blending_mask = (output_max_blending && max_blending_masks) ?
max_blending_masks[pix_id] : true;
typename SplatPrimitive::RenderOutput pix_out = SplatPrimitive::RenderOutput::zero();
typename SplatPrimitive::RenderOutput pix2_out = SplatPrimitive::RenderOutput::zero();
typename SplatPrimitive::RenderOutput distortion_out = SplatPrimitive::RenderOutput::zero();
for (uint32_t b = 0; b < num_batches; ++b) {
// resync all threads before beginning next batch
// end early if entire tile is done
if (__syncthreads_count(done) >= block_size) {
break;
}
// each thread fetch 1 gaussian from front to back
// index of gaussian to load
uint32_t batch_start = range_start + block_size * b;
uint32_t idx = batch_start + tr;
if (idx < range_end) {
int32_t g = flatten_ids[idx]; // flatten index in [I * N] or [nnz]
splat_batch[tr] = SplatPrimitive::Screen::loadWithPrecompute(splat_buffer, g);
if (output_max_blending)
splat_idx_batch[tr] = g;
}
// wait for other threads to collect the gaussians in batch
block.sync();
// process gaussians in the current batch for this pixel
uint32_t batch_size = min(block_size, range_end - batch_start);
for (uint32_t t = 0; (t < batch_size) && !done; ++t) {
typename SplatPrimitive::Screen splat = splat_batch[t];
float alpha = splat.evaluate_alpha(ray_o, ray_d);
if (alpha < ALPHA_THRESHOLD) {
continue;
}
const float next_T = T * (1.0f - alpha);
if (next_T <= 1e-4f) { // this pixel is done: exclusive
done = true;
break;
}
const float vis = alpha * T;
const typename SplatPrimitive::RenderOutput color = splat.evaluate_color(ray_o, ray_d);
if (output_distortion) {
distortion_out += (
color * color * (1.0f - T)
+ color * pix_out * -2.0f
+ pix2_out
) * vis;
pix2_out += color * color * vis;
}
pix_out += color * vis;
cur_idx = batch_start + t;
T = next_T;
if (output_max_blending && out_max_blending != nullptr && max_blending_mask) {
uint32_t splat_idx = splat_idx_batch[t];
atomicMax(out_max_blending + (N == 0 ? splat_idx : splat_idx % N), vis);
}
}
}
if (i < image_height && j < image_width) {
render_Ts[pix_id] = T;
int pix_id_global = image_id * image_height * image_width + pix_id;
// TODO: blend background
pix_out.saveParamsToBuffer(render_colors, pix_id_global);
// index in bin of last gaussian in this pixel
last_ids[pix_id] = static_cast<int32_t>(cur_idx);
// distortion
if (output_distortion) {
pix2_out.saveParamsToBuffer(render_colors2, pix_id_global);
distortion_out.saveParamsToBuffer(render_distortions, pix_id_global);
}
}
}
template <typename SplatPrimitive, gsplat::CameraModelType camera_model, bool output_distortion, bool output_max_blending>
void rasterize_to_pixels_eval3d_fwd_kernel_wrapper(
cudaStream_t stream,
const uint32_t I,
const uint32_t N,
const uint32_t n_isects,
const bool packed,
const typename SplatPrimitive::Screen::Buffer splat_buffer,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const float3 *__restrict__ backgrounds, // [I, 3]
const bool *__restrict__ max_blending_masks, // [B, C, image_width, image_height]
const uint32_t image_width,
const uint32_t image_height,
const uint32_t tile_width,
const uint32_t tile_height,
const int32_t *__restrict__ tile_offsets, // [I, tile_height, tile_width]
const int32_t *__restrict__ flatten_ids, // [n_isects]
typename SplatPrimitive::RenderOutput::Buffer render_colors, // [I, image_height, image_width, ...]
float *__restrict__ render_Ts, // [I, image_height, image_width, 1]
int32_t *__restrict__ last_ids, // [I, image_height, image_width]
typename SplatPrimitive::RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...]
typename SplatPrimitive::RenderOutput::Buffer render_distortions, // [I, image_height, image_width, ...]
float* __restrict__ out_max_blending
) {
// Each block covers a tile on the image. In total there are
// I * tile_height * tile_width blocks.
dim3 threads = {TILE_SIZE, TILE_SIZE, 1};
dim3 grid = {I, tile_height, tile_width};
rasterize_to_pixels_eval3d_fwd_kernel<
SplatPrimitive, camera_model, output_distortion, output_max_blending
><<<grid, threads, 0, stream>>>(
I, N, n_isects, packed,
splat_buffer, viewmats, intrins, dist_coeffs_buffer, backgrounds, max_blending_masks,
image_width, image_height, tile_width, tile_height, tile_offsets, flatten_ids,
render_colors, render_Ts, last_ids,
render_colors2, render_distortions, out_max_blending
);
}
+11 -2
View File
@@ -6,10 +6,11 @@ inline constexpr int TILE_SIZE = 16;
inline constexpr float ALPHA_THRESHOLD = (1.f/255.f);
#ifndef NO_TORCH
#include <c10/cuda/CUDAGuard.h>
#include <ATen/Tensor.h>
#include <ATen/DeviceGuard.h>
#endif
#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be a CUDA tensor")
#define CHECK_HOST(x) TORCH_CHECK(!x.is_cuda(), #x " must be a CPU tensor")
@@ -17,7 +18,7 @@ inline constexpr float ALPHA_THRESHOLD = (1.f/255.f);
TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_INPUT(x) \
do { CHECK_CUDA(x); CHECK_CONTIGUOUS(x); } while (0)
#if 0
#if 1
#define DEVICE_GUARD(_ten) \
const at::cuda::OptionalCUDAGuard device_guard(device_of(_ten));
#else
@@ -69,6 +70,8 @@ inline __host__ dim3 tuple2dim3(std::tuple<unsigned, unsigned, unsigned> v) {
#include "common_utils.cuh"
#ifndef NO_TORCH
template<typename T, int ndim>
TensorView<T, ndim> tensor2view(at::Tensor& tensor) {
TensorView<T, ndim> view;
@@ -80,6 +83,10 @@ TensorView<T, ndim> tensor2view(at::Tensor& tensor) {
return view;
}
#endif
#ifndef NO_TORCH
#include <ATen/ops/empty_like.h>
template<typename T>
@@ -93,3 +100,5 @@ template<typename T>
inline void set_zero(at::Tensor& x) {
cudaMemset(x.data_ptr<T>(), 0, x.numel() * sizeof(T));
}
#endif
+4
View File
@@ -11,7 +11,11 @@
#include "SplatTileIntersector.cuh"
#include "SVHash.cuh"
#include "Projection.cuh"
#include "ProjectionFwd.cuh"
#include "ProjectionBwd.cuh"
#include "Rasterization.cuh"
#include "RasterizationEval3DFwd.cuh"
#include "RasterizationEval3DBwd.cuh"
#include "RasterizationSortedEval3DFwd.cuh"
#include "RasterizationSortedEval3DBwd.cuh"
#include "Optimizer.cuh"
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::AllReasonable
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
MipSplatting::Screen::Buffer v_splats_screen,
MipSplatting::Screen::Buffer vr_splats_screen,
MipSplatting::Screen::Buffer h_splats_screen,
// grad inputs
MipSplatting::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
MipSplatting::World::Buffer vr_splats_world,
MipSplatting::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
MipSplatting::Screen::Buffer v_splats_screen,
MipSplatting::Screen::Buffer vr_splats_screen,
MipSplatting::Screen::Buffer h_splats_screen,
// grad inputs
MipSplatting::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
MipSplatting::World::Buffer vr_splats_world,
MipSplatting::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::Position
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
MipSplatting::Screen::Buffer v_splats_screen,
MipSplatting::Screen::Buffer vr_splats_screen,
MipSplatting::Screen::Buffer h_splats_screen,
// grad inputs
MipSplatting::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
MipSplatting::World::Buffer vr_splats_world,
MipSplatting::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::AllReasonable
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
MipSplatting::Screen::Buffer v_splats_screen,
MipSplatting::Screen::Buffer vr_splats_screen,
MipSplatting::Screen::Buffer h_splats_screen,
// grad inputs
MipSplatting::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
MipSplatting::World::Buffer vr_splats_world,
MipSplatting::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
MipSplatting::Screen::Buffer v_splats_screen,
MipSplatting::Screen::Buffer vr_splats_screen,
MipSplatting::Screen::Buffer h_splats_screen,
// grad inputs
MipSplatting::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
MipSplatting::World::Buffer vr_splats_world,
MipSplatting::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::Position
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
MipSplatting::Screen::Buffer v_splats_screen,
MipSplatting::Screen::Buffer vr_splats_screen,
MipSplatting::Screen::Buffer h_splats_screen,
// grad inputs
MipSplatting::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
MipSplatting::World::Buffer vr_splats_world,
MipSplatting::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "PrimitiveOpaqueTriangle.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
OpaqueTriangle,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const OpaqueTriangle::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
OpaqueTriangle::Screen::Buffer v_splats_screen,
OpaqueTriangle::Screen::Buffer vr_splats_screen,
OpaqueTriangle::Screen::Buffer h_splats_screen,
// grad inputs
OpaqueTriangle::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
OpaqueTriangle::World::Buffer vr_splats_world,
OpaqueTriangle::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "PrimitiveOpaqueTriangle.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
OpaqueTriangle,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const OpaqueTriangle::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
OpaqueTriangle::Screen::Buffer v_splats_screen,
OpaqueTriangle::Screen::Buffer vr_splats_screen,
OpaqueTriangle::Screen::Buffer h_splats_screen,
// grad inputs
OpaqueTriangle::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
OpaqueTriangle::World::Buffer vr_splats_world,
OpaqueTriangle::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<2>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<2>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<2>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<2>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<2>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<2>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<2>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<2>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<2>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<2>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<2>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<2>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<2>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<2>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<2>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<2>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<3>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<3>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<3>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<3>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<3>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<3>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<3>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<3>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<3>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<3>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<3>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<3>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<3>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<3>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<3>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<3>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<4>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<4>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<4>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<4>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<4>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<4>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<4>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<4>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<4>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<4>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<4>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<4>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<4>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<4>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<4>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<4>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<5>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<5>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<5>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<5>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<5>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<5>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<5>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<5>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<5>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<5>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<5>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<5>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<5>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<5>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<5>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<5>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<6>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<6>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<6>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<6>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<6>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<6>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<6>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<6>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<6>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<6>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<6>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<6>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<6>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<6>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<6>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<6>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<7>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<7>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<7>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<7>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<7>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<7>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<7>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<7>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<7>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<7>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<7>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<7>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<7>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<7>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<7>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<7>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<8>,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<8>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<8>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<8>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<8>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<8>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<8>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<8>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
SphericalVoronoi3DGUT<8>,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<8>::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
SphericalVoronoi3DGUT<8>::Screen::Buffer v_splats_screen,
SphericalVoronoi3DGUT<8>::Screen::Buffer vr_splats_screen,
SphericalVoronoi3DGUT<8>::Screen::Buffer h_splats_screen,
// grad inputs
SphericalVoronoi3DGUT<8>::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
SphericalVoronoi3DGUT<8>::World::Buffer vr_splats_world,
SphericalVoronoi3DGUT<8>::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::AllReasonable
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGS::Screen::Buffer v_splats_screen,
Vanilla3DGS::Screen::Buffer vr_splats_screen,
Vanilla3DGS::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGS::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGS::World::Buffer vr_splats_world,
Vanilla3DGS::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGS::Screen::Buffer v_splats_screen,
Vanilla3DGS::Screen::Buffer vr_splats_screen,
Vanilla3DGS::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGS::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGS::World::Buffer vr_splats_world,
Vanilla3DGS::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::Position
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGS::Screen::Buffer v_splats_screen,
Vanilla3DGS::Screen::Buffer vr_splats_screen,
Vanilla3DGS::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGS::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGS::World::Buffer vr_splats_world,
Vanilla3DGS::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::AllReasonable
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGS::Screen::Buffer v_splats_screen,
Vanilla3DGS::Screen::Buffer vr_splats_screen,
Vanilla3DGS::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGS::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGS::World::Buffer vr_splats_world,
Vanilla3DGS::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGS::Screen::Buffer v_splats_screen,
Vanilla3DGS::Screen::Buffer vr_splats_screen,
Vanilla3DGS::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGS::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGS::World::Buffer vr_splats_world,
Vanilla3DGS::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::Position
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGS::Screen::Buffer v_splats_screen,
Vanilla3DGS::Screen::Buffer vr_splats_screen,
Vanilla3DGS::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGS::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGS::World::Buffer vr_splats_world,
Vanilla3DGS::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::AllReasonable
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGUT::Screen::Buffer v_splats_screen,
Vanilla3DGUT::Screen::Buffer vr_splats_screen,
Vanilla3DGUT::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGUT::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGUT::World::Buffer vr_splats_world,
Vanilla3DGUT::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGUT::Screen::Buffer v_splats_screen,
Vanilla3DGUT::Screen::Buffer vr_splats_screen,
Vanilla3DGUT::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGUT::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGUT::World::Buffer vr_splats_world,
Vanilla3DGUT::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::Position
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGUT::Screen::Buffer v_splats_screen,
Vanilla3DGUT::Screen::Buffer vr_splats_screen,
Vanilla3DGUT::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGUT::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGUT::World::Buffer vr_splats_world,
Vanilla3DGUT::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::AllReasonable
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGUT::Screen::Buffer v_splats_screen,
Vanilla3DGUT::Screen::Buffer vr_splats_screen,
Vanilla3DGUT::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGUT::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGUT::World::Buffer vr_splats_world,
Vanilla3DGUT::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGUT::Screen::Buffer v_splats_screen,
Vanilla3DGUT::Screen::Buffer vr_splats_screen,
Vanilla3DGUT::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGUT::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGUT::World::Buffer vr_splats_world,
Vanilla3DGUT::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::Position
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
Vanilla3DGUT::Screen::Buffer v_splats_screen,
Vanilla3DGUT::Screen::Buffer vr_splats_screen,
Vanilla3DGUT::Screen::Buffer h_splats_screen,
// grad inputs
Vanilla3DGUT::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
Vanilla3DGUT::World::Buffer vr_splats_world,
Vanilla3DGUT::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "PrimitiveVoxel.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
VoxelPrimitive,
gsplat::CameraModelType::FISHEYE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const VoxelPrimitive::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
VoxelPrimitive::Screen::Buffer v_splats_screen,
VoxelPrimitive::Screen::Buffer vr_splats_screen,
VoxelPrimitive::Screen::Buffer h_splats_screen,
// grad inputs
VoxelPrimitive::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
VoxelPrimitive::World::Buffer vr_splats_world,
VoxelPrimitive::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,36 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "PrimitiveVoxel.cuh"
#include "ProjectionBwd_kernel.cuh"
template void projection_fused_bwd_kernel_wrapper<
VoxelPrimitive,
gsplat::CameraModelType::PINHOLE,
HessianDiagonalOutputMode::None
>(
cudaStream_t stream,
// fwd inputs
const uint32_t B,
const uint32_t C,
const uint32_t N,
const VoxelPrimitive::World::Buffer splats_world,
const float * viewmats, // [B, C, 4, 4]
const float4 * intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
// fwd outputs
const int4 * aabb, // [B, C, N, 4]
// grad outputs
VoxelPrimitive::Screen::Buffer v_splats_screen,
VoxelPrimitive::Screen::Buffer vr_splats_screen,
VoxelPrimitive::Screen::Buffer h_splats_screen,
// grad inputs
VoxelPrimitive::World::Buffer v_splats_world,
float3* vr_world_pos_buffer,
float3* h_world_pos_buffer,
VoxelPrimitive::World::Buffer vr_splats_world,
VoxelPrimitive::World::Buffer h_splats_world,
float * v_viewmats // [B, C, 4, 4] optional
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
MipSplatting::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
MipSplatting,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const MipSplatting::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
MipSplatting::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "PrimitiveOpaqueTriangle.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
OpaqueTriangle,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const OpaqueTriangle::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
OpaqueTriangle::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "PrimitiveOpaqueTriangle.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
OpaqueTriangle,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const OpaqueTriangle::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
OpaqueTriangle::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<2>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<2>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<2>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<2>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<2>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<2>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<3>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<3>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<3>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<3>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<3>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<3>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<4>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<4>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<4>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<4>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<4>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<4>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<5>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<5>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<5>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<5>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<5>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<5>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<6>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<6>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<6>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<6>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<6>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<6>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<7>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<7>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<7>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<7>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<7>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<7>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<8>,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<8>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<8>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT_SV.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
SphericalVoronoi3DGUT<8>,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const SphericalVoronoi3DGUT<8>::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
SphericalVoronoi3DGUT<8>::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
Vanilla3DGS::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGS.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
Vanilla3DGS,
gsplat::CameraModelType::PINHOLE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGS::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
Vanilla3DGS::Screen::Buffer splats_screen
);
@@ -0,0 +1,26 @@
// This file is auto generated by `generate_kernel_instantiation.py`
#define NO_TORCH
#include "Primitive3DGUT.cuh"
#include "ProjectionFwd_kernel.cuh"
template void projection_fused_fwd_kernel_wrapper<
Vanilla3DGUT,
gsplat::CameraModelType::FISHEYE
>(
cudaStream_t stream,
const uint32_t B,
const uint32_t C,
const uint32_t N,
const Vanilla3DGUT::World::Buffer splats_world,
const float *__restrict__ viewmats, // [B, C, 4, 4]
const float4 *__restrict__ intrins, // [B, C, 4], fx, fy, cx, cy
const CameraDistortionCoeffsBuffer dist_coeffs_buffer,
const uint32_t image_width,
const uint32_t image_height,
const float near_plane,
const float far_plane,
// outputs
int4 *__restrict__ aabbs, // [B, C, N, 4]
Vanilla3DGUT::Screen::Buffer splats_screen
);

Some files were not shown because too many files have changed in this diff Show More