mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
move kernel instantiation to separate files (wip)
This commit is contained in:
@@ -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',
|
||||
|
||||
@@ -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()
|
||||
@@ -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])
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
@@ -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
|
||||
@@ -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"
|
||||
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+36
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
|
||||
);
|
||||
+26
@@ -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
Reference in New Issue
Block a user