""" EndoGaussian-4D Trainer Core training pipeline for deformable 4D Gaussian Splatting in endoscopic scenes. Architecture: G_t = G_0 + Δ_θ(t) G_0: Canonical Gaussians {μ, q, s, α, SH} initialized via Holistic Gaussian Initialization (HGI) from depth backprojection across all frames. Δ_θ: HexPlane-encoded deformation field - Encoder: 6 feature planes (XY,XZ,YZ,XT,YT,ZT) with bilinear sampling - Decoder: Shared MLP → 4 heads (Δμ, Δq, Δs, Δα), zero-initialized Loss: L = L_rgb + λ₁·L_dssim + λ₂·L_depth + λ₃·L_smooth + λ₄·L_tv Training recipe from EndoGaussian (Liu et al. 2024): - 3000 iterations total, 1000 warmup (static only, no deformation) - Adam optimizer, lr_means=1.6e-4 with exponential decay - HexPlane resolution: 64³ spatial × 75 temporal - Densification every 100 steps via absgrad (gsplat) - Tool masking in both initialization and loss computation Reference: gsplat (nerfstudio-project/gsplat) for differentiable rasterization """ import json import math import os import time from dataclasses import dataclass, field from pathlib import Path from typing import Dict, List, Optional, Tuple import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from torch.optim import Adam from torch.optim.lr_scheduler import ExponentialLR from .metrics import compute_psnr, compute_ssim, dssim_loss, LPIPSMetric, MetricsAccumulator # --------------------------------------------------------------------------- # Configuration # --------------------------------------------------------------------------- @dataclass class EndoGaussianConfig: """Full hyperparameter configuration for EndoGaussian-4D training.""" # --- Training schedule --- total_iters: int = 3000 # Total training iterations warmup_iters: int = 1000 # Static-only warmup (no deformation) densify_start: int = 500 # Start densification densify_stop: int = 2500 # Stop densification densify_interval: int = 100 # Densify every N steps prune_interval: int = 100 # Prune every N steps eval_interval: int = 500 # Evaluate every N steps checkpoint_interval: int = 1000 # Save checkpoint every N steps # --- Learning rates --- lr_means: float = 1.6e-4 # Position learning rate lr_scales: float = 5e-3 # Scale learning rate lr_quats: float = 1e-3 # Rotation learning rate lr_opacities: float = 5e-2 # Opacity learning rate lr_sh: float = 2.5e-3 # Spherical harmonics learning rate lr_deformation: float = 1.6e-3 # Deformation network learning rate lr_decay_factor: float = 0.01 # Final lr = initial * decay_factor lr_decay_steps: int = 3000 # Steps for full decay # --- Loss weights --- lambda_dssim: float = 0.2 # D-SSIM weight lambda_depth: float = 0.1 # Depth supervision weight lambda_smooth: float = 0.01 # Temporal smoothness weight lambda_tv: float = 0.001 # Total variation on HexPlane lambda_consistency: float = 0.0001 # Cross-frame consistency # --- Densification thresholds --- densify_grad_thresh: float = 0.0002 # Gradient threshold for split/clone densify_size_thresh: float = 0.01 # Size threshold for splitting prune_opacity_thresh: float = 0.005 # Prune Gaussians below this opacity prune_size_thresh: float = 0.1 # Prune Gaussians above this size max_gaussians: int = 500_000 # Maximum Gaussian count # --- HexPlane deformation --- hexplane_resolution: List[int] = field(default_factory=lambda: [64, 64, 64, 75]) hexplane_num_levels: int = 2 # Multi-resolution levels hexplane_feat_dim: int = 32 # Feature dimension per plane deform_hidden_dim: int = 128 # MLP hidden dimension deform_num_layers: int = 3 # MLP depth sh_degree: int = 3 # Spherical harmonics degree # --- Data --- image_height: int = 540 image_width: int = 675 train_ratio: float = 0.875 # 7:1 train/test split # --- Device --- device: str = "cuda" # --- Output --- output_dir: str = "./output" experiment_name: str = "endogaussian4d" # --------------------------------------------------------------------------- # HexPlane Temporal Encoder # --------------------------------------------------------------------------- class HexPlaneEncoder(nn.Module): """ HexPlane feature encoding for spatio-temporal deformation. Factorizes 4D (x,y,z,t) space into 6 feature planes: Spatial: XY, XZ, YZ Temporal: XT, YT, ZT Each plane stores learned features at multiple resolutions. Features are extracted via bilinear interpolation and concatenated. The factorization reduces memory from O(N⁴) to O(6·N²), enabling real-time deformation of 100K+ Gaussians. Args: resolution: [Rx, Ry, Rz, Rt] grid resolution num_levels: Number of multi-resolution levels feat_dim: Feature dimension per plane per level """ # The 6 planes and which coordinate indices they use PLANE_AXES = [ (0, 1), # XY - spatial (0, 2), # XZ - spatial (1, 2), # YZ - spatial (0, 3), # XT - temporal (1, 3), # YT - temporal (2, 3), # ZT - temporal ] def __init__(self, resolution: List[int], num_levels: int = 2, feat_dim: int = 32): super().__init__() self.resolution = resolution # [Rx, Ry, Rz, Rt] self.num_levels = num_levels self.feat_dim = feat_dim self.output_dim = 6 * num_levels * feat_dim # Create learnable feature planes at each resolution level self.planes = nn.ParameterList() for level in range(num_levels): scale = 2 ** level for ax_i, ax_j in self.PLANE_AXES: res_i = resolution[ax_i] * scale res_j = resolution[ax_j] * scale # Initialize with small random values plane = nn.Parameter(0.1 * torch.randn(1, feat_dim, res_i, res_j)) self.planes.append(plane) def forward(self, coords: torch.Tensor) -> torch.Tensor: """ Sample features from all planes. Args: coords: [N, 4] tensor of (x, y, z, t) coordinates, all in [-1, 1] Returns: [N, output_dim] concatenated features from all planes """ batch_size = coords.shape[0] features = [] plane_idx = 0 for level in range(self.num_levels): for ax_i, ax_j in self.PLANE_AXES: plane = self.planes[plane_idx] # Extract the 2 coordinates for this plane grid_coords = coords[:, [ax_i, ax_j]] # [N, 2] # Reshape for grid_sample: [1, 1, N, 2] grid = grid_coords.view(1, 1, batch_size, 2) # Bilinear interpolation sampled = F.grid_sample( plane, grid, mode="bilinear", padding_mode="border", align_corners=True, ) # [1, feat_dim, 1, N] features.append(sampled.squeeze(0).squeeze(1).T) # [N, feat_dim] plane_idx += 1 return torch.cat(features, dim=-1) # [N, 6 * num_levels * feat_dim] def tv_loss(self) -> torch.Tensor: """ Total variation regularization on feature planes. Encourages spatial smoothness in the learned features, preventing noisy deformation artifacts. Critical for temporal planes (XT, YT, ZT) to ensure smooth motion over time. """ loss = torch.tensor(0.0, device=self.planes[0].device) for plane in self.planes: # Horizontal TV loss = loss + (plane[:, :, :, 1:] - plane[:, :, :, :-1]).abs().mean() # Vertical TV loss = loss + (plane[:, :, 1:, :] - plane[:, :, :-1, :]).abs().mean() return loss / len(self.planes) # --------------------------------------------------------------------------- # Deformation Decoder # --------------------------------------------------------------------------- class DeformationDecoder(nn.Module): """ Decodes HexPlane features into Gaussian parameter deltas. Architecture: Shared MLP backbone → 4 prediction heads: - Δμ: Position displacement [N, 3] - Δq: Rotation perturbation [N, 4] (added to canonical quaternion) - Δs: Scale adjustment [N, 3] - Δα: Opacity adjustment [N, 1] The heads are ZERO-INITIALIZED, which is critical for stable training. At initialization, the deformation is identity (Δ = 0), so the model starts from the canonical Gaussians and gradually learns displacements. Args: input_dim: Feature dimension from HexPlane encoder hidden_dim: Hidden layer dimension num_layers: Number of hidden layers in shared backbone """ def __init__(self, input_dim: int, hidden_dim: int = 128, num_layers: int = 3): super().__init__() # Shared backbone layers = [] in_dim = input_dim for _ in range(num_layers): layers.extend([ nn.Linear(in_dim, hidden_dim), nn.ReLU(inplace=True), ]) in_dim = hidden_dim self.backbone = nn.Sequential(*layers) # Prediction heads — zero initialized! self.head_means = nn.Linear(hidden_dim, 3) self.head_quats = nn.Linear(hidden_dim, 4) self.head_scales = nn.Linear(hidden_dim, 3) self.head_opacities = nn.Linear(hidden_dim, 1) # Zero-initialize all heads for identity deformation at start for head in [self.head_means, self.head_quats, self.head_scales, self.head_opacities]: nn.init.zeros_(head.weight) nn.init.zeros_(head.bias) def forward(self, features: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """ Predict deformation deltas from HexPlane features. Args: features: [N, input_dim] from HexPlane encoder Returns: Tuple of (delta_means, delta_quats, delta_scales, delta_opacities) Each [N, D] where D is the parameter dimension. """ h = self.backbone(features) return ( self.head_means(h), # [N, 3] self.head_quats(h), # [N, 4] self.head_scales(h), # [N, 3] self.head_opacities(h), # [N, 1] ) # --------------------------------------------------------------------------- # Depth Loss (Scale-Invariant) # --------------------------------------------------------------------------- def scale_invariant_depth_loss( pred_depth: torch.Tensor, gt_depth: torch.Tensor, mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Scale-invariant depth loss (Eigen et al. 2014). Since endoscopic depth is often relative (from monocular estimation or structured-light with unknown offset), we use a scale-invariant formulation that's robust to global scale/shift ambiguity: L_depth = (1/n) Σ d_i² - (λ/n²)(Σ d_i)² where d_i = log(pred_i) - log(gt_i), λ = 0.5 Args: pred_depth: [B, 1, H, W] predicted depth gt_depth: [B, 1, H, W] ground truth depth mask: [B, 1, H, W] optional validity mask """ if mask is None: mask = (gt_depth > 1e-6).float() pred_log = torch.log(torch.clamp(pred_depth, min=1e-6)) gt_log = torch.log(torch.clamp(gt_depth, min=1e-6)) diff = (pred_log - gt_log) * mask n = mask.sum() + 1e-8 loss = (diff ** 2).sum() / n - 0.5 * (diff.sum() ** 2) / (n ** 2) return loss # --------------------------------------------------------------------------- # EndoGaussian Trainer # --------------------------------------------------------------------------- class EndoGaussianTrainer: """ Main trainer for EndoGaussian-4D. Manages the full pipeline: 1. Point cloud initialization (from HGI depth backprojection) 2. Gaussian parameter optimization with adaptive density control 3. HexPlane-encoded deformation field training 4. Evaluation with PSNR/SSIM/LPIPS/depth metrics 5. Checkpointing and PLY export The training follows a two-phase schedule: Phase 1 (warmup, iters 0-1000): Train static Gaussians only. The deformation network exists but outputs zeros. Phase 2 (deformation, iters 1000-3000): Joint optimization of Gaussians + deformation field. Usage: config = EndoGaussianConfig() trainer = EndoGaussianTrainer(config) trainer.initialize_from_point_cloud(points, colors, cameras) trainer.train(dataset) """ def __init__(self, config: EndoGaussianConfig): self.config = config self.device = torch.device(config.device) self.step = 0 # --- Gaussian parameters (canonical frame) --- self.means: Optional[torch.Tensor] = None # [N, 3] positions self.quats: Optional[torch.Tensor] = None # [N, 4] quaternions (wxyz) self.scales: Optional[torch.Tensor] = None # [N, 3] log-scales self.opacities: Optional[torch.Tensor] = None # [N, 1] logit-opacities self.sh_coeffs: Optional[torch.Tensor] = None # [N, K, 3] SH coefficients # --- Deformation networks --- self.hexplane = HexPlaneEncoder( resolution=config.hexplane_resolution, num_levels=config.hexplane_num_levels, feat_dim=config.hexplane_feat_dim, ).to(self.device) hexplane_out_dim = self.hexplane.output_dim self.decoder = DeformationDecoder( input_dim=hexplane_out_dim, hidden_dim=config.deform_hidden_dim, num_layers=config.deform_num_layers, ).to(self.device) # --- Scene bounds (set during initialization) --- self.scene_center = torch.zeros(3, device=self.device) self.scene_scale = torch.ones(3, device=self.device) self.time_min = 0.0 self.time_max = 1.0 # --- Optimizer (built after initialization) --- self.optimizer: Optional[Adam] = None self.scheduler: Optional[ExponentialLR] = None # --- Training state --- self.grad_accum: Optional[torch.Tensor] = None # For absgrad densification self.grad_count: Optional[torch.Tensor] = None self.max_radii: Optional[torch.Tensor] = None @property def num_gaussians(self) -> int: return self.means.shape[0] if self.means is not None else 0 # ------------------------------------------------------------------- # Initialization # ------------------------------------------------------------------- def initialize_from_point_cloud( self, points: np.ndarray, colors: np.ndarray, cameras: List[Dict], subsample_ratio: float = 0.001, ): """ Initialize Gaussians from a point cloud (HGI output). Implements Holistic Gaussian Initialization: P = ∪_t K⁻¹ · T_t · D_t · (I_t ⊙ M_t) The union of depth-backprojected points across all frames provides complete scene coverage, avoiding the sparse-initialization problem of vanilla 3DGS on endoscopic scenes. Args: points: [N, 3] world-space point cloud colors: [N, 3] RGB colors in [0, 1] cameras: List of camera dicts with 'extrinsic', 'intrinsic' keys subsample_ratio: Fraction of points to keep (0.001 = 0.1%) """ # Subsample if too many points n_points = points.shape[0] if subsample_ratio < 1.0: n_keep = max(int(n_points * subsample_ratio), 1000) indices = np.random.choice(n_points, n_keep, replace=False) points = points[indices] colors = colors[indices] N = points.shape[0] print(f"[Init] Initializing {N:,} Gaussians from point cloud") # Compute scene bounds for coordinate normalization self.scene_center = torch.from_numpy(points.mean(axis=0)).float().to(self.device) extent = points.max(axis=0) - points.min(axis=0) self.scene_scale = torch.from_numpy(extent).float().to(self.device).clamp(min=1e-6) # --- Means --- self.means = torch.from_numpy(points).float().to(self.device) self.means.requires_grad_(True) # --- Scales: estimate from k-nearest-neighbor distances --- # Use a simple distance-based heuristic: scale ~ mean distance to 3 nearest neighbors try: from scipy.spatial import KDTree tree = KDTree(points) dists, _ = tree.query(points, k=4) # k=4 because first neighbor is self nn_dist = np.mean(dists[:, 1:], axis=1) # Mean of 3 nearest log_scales = np.log(np.clip(nn_dist, 1e-7, None)) except ImportError: # Fallback: uniform scale log_scales = np.full(N, np.log(0.001)) self.scales = torch.from_numpy( np.stack([log_scales] * 3, axis=-1) ).float().to(self.device) self.scales.requires_grad_(True) # --- Rotations: identity quaternions --- self.quats = torch.zeros(N, 4, device=self.device) self.quats[:, 0] = 1.0 # w=1, xyz=0 → identity rotation self.quats.requires_grad_(True) # --- Opacities: sigmoid⁻¹(0.1) → starts mostly transparent --- init_opacity = 0.1 logit_opacity = math.log(init_opacity / (1.0 - init_opacity)) self.opacities = torch.full((N, 1), logit_opacity, device=self.device) self.opacities.requires_grad_(True) # --- SH coefficients: DC term from RGB, higher orders zero --- num_sh = (self.config.sh_degree + 1) ** 2 self.sh_coeffs = torch.zeros(N, num_sh, 3, device=self.device) # Set DC (0th order) from colors: C0 = (color - 0.5) / SH_C0 SH_C0 = 0.28209479177387814 # 1 / (2*sqrt(pi)) self.sh_coeffs[:, 0, :] = ( torch.from_numpy(colors).float().to(self.device) - 0.5 ) / SH_C0 self.sh_coeffs.requires_grad_(True) # --- Gradient accumulators for densification --- self.grad_accum = torch.zeros(N, device=self.device) self.grad_count = torch.zeros(N, device=self.device, dtype=torch.int32) self.max_radii = torch.zeros(N, device=self.device) # --- Build optimizer --- self._build_optimizer() print(f"[Init] Done. Scene center: {self.scene_center.cpu().numpy()}, " f"scale: {self.scene_scale.cpu().numpy()}") def _build_optimizer(self): """Construct Adam optimizer with per-parameter-group learning rates.""" cfg = self.config param_groups = [ {"params": [self.means], "lr": cfg.lr_means, "name": "means"}, {"params": [self.scales], "lr": cfg.lr_scales, "name": "scales"}, {"params": [self.quats], "lr": cfg.lr_quats, "name": "quats"}, {"params": [self.opacities], "lr": cfg.lr_opacities, "name": "opacities"}, {"params": [self.sh_coeffs], "lr": cfg.lr_sh, "name": "sh_coeffs"}, {"params": self.hexplane.parameters(), "lr": cfg.lr_deformation, "name": "hexplane"}, {"params": self.decoder.parameters(), "lr": cfg.lr_deformation, "name": "decoder"}, ] self.optimizer = Adam(param_groups, eps=1e-15) # Exponential decay: lr(t) = lr_0 * (decay_factor)^(t/decay_steps) gamma = cfg.lr_decay_factor ** (1.0 / cfg.lr_decay_steps) self.scheduler = ExponentialLR(self.optimizer, gamma=gamma) # ------------------------------------------------------------------- # Deformation # ------------------------------------------------------------------- def apply_deformation( self, timestamp: float, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """ Apply temporal deformation to canonical Gaussians. G_t = G_0 + Δ_θ(t) During warmup (step < warmup_iters), returns canonical params unchanged to let the static Gaussians converge first. Args: timestamp: Normalized time in [0, 1] Returns: Tuple of (deformed_means, deformed_quats, deformed_scales, deformed_opacities) """ if self.step < self.config.warmup_iters: # Warmup phase: no deformation return self.means, self.quats, self.scales, self.opacities # Normalize positions to [-1, 1] for grid sampling norm_pos = (self.means - self.scene_center) / (self.scene_scale / 2.0) norm_pos = torch.clamp(norm_pos, -1.0, 1.0) # Normalize time to [-1, 1] norm_t = 2.0 * timestamp - 1.0 time_col = torch.full((self.num_gaussians, 1), norm_t, device=self.device) # 4D coordinates: [N, 4] = [x, y, z, t] coords = torch.cat([norm_pos, time_col], dim=-1) # Encode via HexPlane features = self.hexplane(coords) # [N, feat_dim] # Decode deltas d_means, d_quats, d_scales, d_opacs = self.decoder(features) # Apply deformations deformed_means = self.means + d_means deformed_quats = self.quats + d_quats deformed_scales = self.scales + d_scales deformed_opacities = self.opacities + d_opacs return deformed_means, deformed_quats, deformed_scales, deformed_opacities # ------------------------------------------------------------------- # Rendering # ------------------------------------------------------------------- def render( self, viewmat: torch.Tensor, K: torch.Tensor, width: int, height: int, timestamp: float = 0.0, near: float = 0.01, far: float = 100.0, ) -> Dict[str, torch.Tensor]: """ Render a frame at the given camera pose and timestamp. Uses gsplat.rasterization for differentiable splatting with depth output for supervision. Args: viewmat: [4, 4] world-to-camera transform K: [3, 3] intrinsic matrix width, height: Image dimensions timestamp: Normalized time [0, 1] for deformation near, far: Clipping planes Returns: Dict with keys: "rgb": [H, W, 3] rendered color "depth": [H, W, 1] rendered depth "alpha": [H, W, 1] rendered alpha (opacity) """ try: from gsplat import rasterization except ImportError: raise ImportError("gsplat not installed. Run: pip install gsplat>=1.4.0") # Get deformed parameters for this timestamp d_means, d_quats, d_scales, d_opacs = self.apply_deformation(timestamp) # Activate parameters activated_scales = torch.exp(d_scales) # log-scale → scale activated_opacities = torch.sigmoid(d_opacs) # logit → probability quats_normalized = F.normalize(d_quats, dim=-1) # Normalize quaternions # Reshape for gsplat batch dimension viewmat_4x4 = viewmat.unsqueeze(0).to(self.device) # [1, 4, 4] K_3x3 = K.unsqueeze(0).to(self.device) # [1, 3, 3] # gsplat rasterization renders, alphas, meta = rasterization( means=d_means, quats=quats_normalized, scales=activated_scales, opacities=activated_opacities.squeeze(-1), colors=self.sh_coeffs, viewmats=viewmat_4x4, Ks=K_3x3, width=width, height=height, near_plane=near, far_plane=far, sh_degree=self.config.sh_degree, render_mode="RGB+D", absgrad=True, # Enable absolute gradient for better densification ) # renders: [1, H, W, 4] (RGB+D) rgb = renders[0, :, :, :3] # [H, W, 3] depth = renders[0, :, :, 3:4] # [H, W, 1] alpha = alphas[0, :, :, None] # [H, W, 1] return { "rgb": rgb, "depth": depth, "alpha": alpha, "meta": meta, } # ------------------------------------------------------------------- # Loss computation # ------------------------------------------------------------------- def compute_loss( self, rendered: Dict[str, torch.Tensor], gt_rgb: torch.Tensor, gt_depth: Optional[torch.Tensor] = None, tool_mask: Optional[torch.Tensor] = None, prev_rendered: Optional[Dict[str, torch.Tensor]] = None, ) -> Tuple[torch.Tensor, Dict[str, float]]: """ Compute the full EndoGaussian-4D loss. L = L_rgb + λ₁·L_dssim + λ₂·L_depth + λ₃·L_smooth + λ₄·L_tv Tool masks exclude surgical instruments from loss computation since they are rigid objects that shouldn't be modeled as tissue. Args: rendered: Output from self.render() gt_rgb: [H, W, 3] ground truth image gt_depth: [H, W, 1] ground truth depth (optional) tool_mask: [H, W, 1] binary mask (1=tissue, 0=tool) prev_rendered: Previous frame render for temporal smoothness Returns: (total_loss, loss_dict) where loss_dict has individual terms """ cfg = self.config pred_rgb = rendered["rgb"] pred_depth = rendered["depth"] loss_dict = {} # Apply tool mask if available if tool_mask is not None: mask = tool_mask.float() pred_rgb_masked = pred_rgb * mask gt_rgb_masked = gt_rgb * mask else: mask = None pred_rgb_masked = pred_rgb gt_rgb_masked = gt_rgb # --- L1 RGB loss --- l1_loss = F.l1_loss(pred_rgb_masked, gt_rgb_masked) loss_dict["l1"] = l1_loss.item() # --- D-SSIM loss --- # Reshape to [1, 3, H, W] for SSIM computation pred_4d = pred_rgb_masked.permute(2, 0, 1).unsqueeze(0) gt_4d = gt_rgb_masked.permute(2, 0, 1).unsqueeze(0) dssim = dssim_loss(pred_4d, gt_4d) loss_dict["dssim"] = dssim.item() total_loss = (1.0 - cfg.lambda_dssim) * l1_loss + cfg.lambda_dssim * dssim # --- Depth loss --- if gt_depth is not None and cfg.lambda_depth > 0: depth_mask = mask if mask is not None else (gt_depth > 1e-6).float() d_loss = scale_invariant_depth_loss( pred_depth.unsqueeze(0).permute(0, 3, 1, 2), gt_depth.unsqueeze(0).permute(0, 3, 1, 2), depth_mask.unsqueeze(0).permute(0, 3, 1, 2) if depth_mask.dim() == 3 else None, ) total_loss = total_loss + cfg.lambda_depth * d_loss loss_dict["depth"] = d_loss.item() # --- Temporal smoothness loss --- if prev_rendered is not None and cfg.lambda_smooth > 0: smooth_loss = F.mse_loss(rendered["rgb"], prev_rendered["rgb"]) total_loss = total_loss + cfg.lambda_smooth * smooth_loss loss_dict["smooth"] = smooth_loss.item() # --- HexPlane TV regularization --- if cfg.lambda_tv > 0 and self.step >= cfg.warmup_iters: tv_loss = self.hexplane.tv_loss() total_loss = total_loss + cfg.lambda_tv * tv_loss loss_dict["tv"] = tv_loss.item() loss_dict["total"] = total_loss.item() return total_loss, loss_dict # ------------------------------------------------------------------- # Adaptive Density Control # ------------------------------------------------------------------- def _update_densification_stats(self, meta: dict): """Track gradient statistics for absgrad-based densification.""" if "means2d" in meta and meta["means2d"].grad is not None: grads = meta["means2d"].grad.detach() # absgrad: use absolute value instead of norm grad_norms = grads.abs().max(dim=-1).values visible = meta.get("gaussian_ids", torch.arange(self.num_gaussians, device=self.device)) if visible.max() < self.num_gaussians: self.grad_accum[visible] += grad_norms self.grad_count[visible] += 1 def densify_and_prune(self): """ Adaptive density control: split, clone, and prune Gaussians. Split: Large Gaussians with high gradient → split into 2 smaller ones Clone: Small Gaussians with high gradient → duplicate at same position Prune: Gaussians with very low opacity or very large scale → remove Based on absgrad statistics accumulated over the densification interval. """ cfg = self.config if self.grad_count is None or (self.grad_count == 0).all(): return # Average gradient over accumulation window avg_grad = self.grad_accum / self.grad_count.clamp(min=1).float() # Identify candidates high_grad_mask = avg_grad > cfg.densify_grad_thresh activated_scales = torch.exp(self.scales) large_mask = activated_scales.max(dim=-1).values > cfg.densify_size_thresh small_mask = ~large_mask # --- Clone: small Gaussians with high gradient --- clone_mask = high_grad_mask & small_mask if clone_mask.any() and self.num_gaussians < cfg.max_gaussians: n_clone = min(clone_mask.sum().item(), cfg.max_gaussians - self.num_gaussians) clone_indices = clone_mask.nonzero(as_tuple=True)[0][:n_clone] self._clone_gaussians(clone_indices) # --- Split: large Gaussians with high gradient --- split_mask = high_grad_mask & large_mask if split_mask.any(): n_split = min(split_mask.sum().item(), cfg.max_gaussians - self.num_gaussians) split_indices = split_mask.nonzero(as_tuple=True)[0][:n_split] self._split_gaussians(split_indices) # --- Prune: low opacity or oversized --- with torch.no_grad(): opacity_vals = torch.sigmoid(self.opacities).squeeze(-1) prune_mask = opacity_vals < cfg.prune_opacity_thresh if cfg.prune_size_thresh > 0: prune_mask = prune_mask | (activated_scales.max(dim=-1).values > cfg.prune_size_thresh) if prune_mask.any(): keep_mask = ~prune_mask self._prune_gaussians(keep_mask) # Reset gradient accumulators self.grad_accum = torch.zeros(self.num_gaussians, device=self.device) self.grad_count = torch.zeros(self.num_gaussians, device=self.device, dtype=torch.int32) self.max_radii = torch.zeros(self.num_gaussians, device=self.device) def _clone_gaussians(self, indices: torch.Tensor): """Clone selected Gaussians (duplicate at same position).""" with torch.no_grad(): new_means = self.means[indices].clone() new_quats = self.quats[indices].clone() new_scales = self.scales[indices].clone() new_opacities = self.opacities[indices].clone() new_sh = self.sh_coeffs[indices].clone() self._append_gaussians(new_means, new_quats, new_scales, new_opacities, new_sh) def _split_gaussians(self, indices: torch.Tensor): """Split selected Gaussians into 2 smaller ones.""" with torch.no_grad(): # Sample offsets from the Gaussian's own distribution scales = torch.exp(self.scales[indices]) stds = scales.detach() offsets = torch.randn_like(stds) * stds new_means = self.means[indices].clone() + offsets new_quats = self.quats[indices].clone() # Reduce scale by factor of 1.6 (preserving volume approximately) new_scales = self.scales[indices].clone() - math.log(1.6) new_opacities = self.opacities[indices].clone() new_sh = self.sh_coeffs[indices].clone() # Also shrink the original Gaussians self.scales.data[indices] -= math.log(1.6) self._append_gaussians(new_means, new_quats, new_scales, new_opacities, new_sh) def _append_gaussians(self, means, quats, scales, opacities, sh_coeffs): """Append new Gaussians and rebuild optimizer.""" # Detach and re-enable gradients self.means = nn.Parameter(torch.cat([self.means.data, means], dim=0)) self.quats = nn.Parameter(torch.cat([self.quats.data, quats], dim=0)) self.scales = nn.Parameter(torch.cat([self.scales.data, scales], dim=0)) self.opacities = nn.Parameter(torch.cat([self.opacities.data, opacities], dim=0)) self.sh_coeffs = nn.Parameter(torch.cat([self.sh_coeffs.data, sh_coeffs], dim=0)) # Extend gradient accumulators n_new = means.shape[0] self.grad_accum = torch.cat([self.grad_accum, torch.zeros(n_new, device=self.device)]) self.grad_count = torch.cat([self.grad_count, torch.zeros(n_new, device=self.device, dtype=torch.int32)]) self.max_radii = torch.cat([self.max_radii, torch.zeros(n_new, device=self.device)]) self._build_optimizer() def _prune_gaussians(self, keep_mask: torch.Tensor): """Remove Gaussians where keep_mask is False.""" self.means = nn.Parameter(self.means.data[keep_mask]) self.quats = nn.Parameter(self.quats.data[keep_mask]) self.scales = nn.Parameter(self.scales.data[keep_mask]) self.opacities = nn.Parameter(self.opacities.data[keep_mask]) self.sh_coeffs = nn.Parameter(self.sh_coeffs.data[keep_mask]) self.grad_accum = self.grad_accum[keep_mask] self.grad_count = self.grad_count[keep_mask] self.max_radii = self.max_radii[keep_mask] self._build_optimizer() # ------------------------------------------------------------------- # Training Loop # ------------------------------------------------------------------- def train(self, dataset, output_dir: Optional[str] = None): """ Full training loop. Args: dataset: Object providing: - __len__(): number of frames - __getitem__(idx): dict with 'rgb' [H,W,3], 'depth' [H,W,1], 'mask' [H,W,1], 'viewmat' [4,4], 'K' [3,3], 'timestamp' float output_dir: Override output directory """ cfg = self.config out_dir = Path(output_dir or cfg.output_dir) / cfg.experiment_name out_dir.mkdir(parents=True, exist_ok=True) # Save config with open(out_dir / "config.json", "w") as f: json.dump(vars(cfg), f, indent=2, default=str) # Train/test split n_frames = len(dataset) n_train = int(n_frames * cfg.train_ratio) all_indices = list(range(n_frames)) train_indices = all_indices[:n_train] test_indices = all_indices[n_train:] print(f"\n{'='*60}") print(f"EndoGaussian-4D Training") print(f"{'='*60}") print(f"Total frames: {n_frames} (train: {n_train}, test: {len(test_indices)})") print(f"Gaussians: {self.num_gaussians:,}") print(f"Iterations: {cfg.total_iters} (warmup: {cfg.warmup_iters})") print(f"Output: {out_dir}") print(f"{'='*60}\n") metrics_log = [] start_time = time.time() for self.step in range(cfg.total_iters): self.optimizer.zero_grad() # Random training frame idx = train_indices[np.random.randint(0, n_train)] sample = dataset[idx] gt_rgb = sample["rgb"].to(self.device) # [H, W, 3] viewmat = sample["viewmat"].to(self.device) # [4, 4] K = sample["K"].to(self.device) # [3, 3] timestamp = sample.get("timestamp", idx / max(n_frames - 1, 1)) gt_depth = sample.get("depth") tool_mask = sample.get("mask") if gt_depth is not None: gt_depth = gt_depth.to(self.device) if tool_mask is not None: tool_mask = tool_mask.to(self.device) # Render rendered = self.render( viewmat=viewmat, K=K, width=cfg.image_width, height=cfg.image_height, timestamp=timestamp, ) # Loss loss, loss_dict = self.compute_loss( rendered=rendered, gt_rgb=gt_rgb, gt_depth=gt_depth, tool_mask=tool_mask, ) # Backward loss.backward() # Update densification stats if rendered.get("meta"): self._update_densification_stats(rendered["meta"]) # Optimizer step self.optimizer.step() self.scheduler.step() # --- Adaptive density control --- if (cfg.densify_start <= self.step < cfg.densify_stop and self.step % cfg.densify_interval == 0): self.densify_and_prune() # --- Logging --- if self.step % 50 == 0: elapsed = time.time() - start_time phase = "warmup" if self.step < cfg.warmup_iters else "deform" loss_str = " | ".join(f"{k}: {v:.4f}" for k, v in loss_dict.items()) print(f"[Step {self.step:5d}/{cfg.total_iters}] [{phase}] " f"{loss_str} | #G: {self.num_gaussians:,} | " f"Time: {elapsed:.1f}s") # --- Evaluation --- if self.step > 0 and self.step % cfg.eval_interval == 0 and test_indices: eval_results = self._evaluate(dataset, test_indices) metrics_log.append({"step": self.step, **eval_results}) print(f"\n [EVAL] Step {self.step}: " f"PSNR={eval_results.get('psnr', 0):.2f} dB | " f"SSIM={eval_results.get('ssim', 0):.4f}\n") # --- Checkpoint --- if self.step > 0 and self.step % cfg.checkpoint_interval == 0: self.save_checkpoint(out_dir / f"ckpt_{self.step:06d}.pth") # Final save self.save_checkpoint(out_dir / "ckpt_final.pth") self.export_ply(out_dir / "gaussians_final.ply", timestamp=0.5) # Save metrics log with open(out_dir / "metrics_log.json", "w") as f: json.dump(metrics_log, f, indent=2) total_time = time.time() - start_time print(f"\n{'='*60}") print(f"Training complete in {total_time:.1f}s ({total_time/60:.1f} min)") print(f"Final Gaussians: {self.num_gaussians:,}") print(f"Output saved to: {out_dir}") print(f"{'='*60}") def _evaluate(self, dataset, test_indices: List[int]) -> Dict[str, float]: """Run evaluation on test set.""" cfg = self.config acc = MetricsAccumulator(device=str(self.device), compute_lpips=False, compute_depth=True) self.hexplane.eval() self.decoder.eval() with torch.no_grad(): for idx in test_indices[:8]: # Eval on subset for speed sample = dataset[idx] gt_rgb = sample["rgb"].to(self.device) viewmat = sample["viewmat"].to(self.device) K = sample["K"].to(self.device) timestamp = sample.get("timestamp", idx / max(len(dataset) - 1, 1)) rendered = self.render( viewmat=viewmat, K=K, width=cfg.image_width, height=cfg.image_height, timestamp=timestamp, ) # Reshape for metrics: [1, 3, H, W] pred_4d = rendered["rgb"].permute(2, 0, 1).unsqueeze(0).clamp(0, 1) gt_4d = gt_rgb.permute(2, 0, 1).unsqueeze(0).clamp(0, 1) pred_d = rendered["depth"].permute(2, 0, 1).unsqueeze(0) if "depth" in rendered else None gt_d = sample.get("depth") if gt_d is not None: gt_d = gt_d.to(self.device).permute(2, 0, 1).unsqueeze(0) if gt_d.dim() == 3 else None acc.update(pred_4d, gt_4d, pred_depth=pred_d, gt_depth=gt_d) self.hexplane.train() self.decoder.train() return acc.compute() # ------------------------------------------------------------------- # Checkpointing # ------------------------------------------------------------------- def save_checkpoint(self, path: str): """Save full training state.""" path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) state = { "step": self.step, "means": self.means.data.cpu(), "quats": self.quats.data.cpu(), "scales": self.scales.data.cpu(), "opacities": self.opacities.data.cpu(), "sh_coeffs": self.sh_coeffs.data.cpu(), "hexplane_state": self.hexplane.state_dict(), "decoder_state": self.decoder.state_dict(), "scene_center": self.scene_center.cpu(), "scene_scale": self.scene_scale.cpu(), "config": vars(self.config), } if self.optimizer is not None: state["optimizer_state"] = self.optimizer.state_dict() torch.save(state, path) print(f"[Checkpoint] Saved to {path}") def load_checkpoint(self, path: str): """Load training state from checkpoint.""" state = torch.load(path, map_location=self.device) self.step = state["step"] self.means = nn.Parameter(state["means"].to(self.device)) self.quats = nn.Parameter(state["quats"].to(self.device)) self.scales = nn.Parameter(state["scales"].to(self.device)) self.opacities = nn.Parameter(state["opacities"].to(self.device)) self.sh_coeffs = nn.Parameter(state["sh_coeffs"].to(self.device)) self.scene_center = state["scene_center"].to(self.device) self.scene_scale = state["scene_scale"].to(self.device) self.hexplane.load_state_dict(state["hexplane_state"]) self.decoder.load_state_dict(state["decoder_state"]) # Reinit gradient accumulators N = self.num_gaussians self.grad_accum = torch.zeros(N, device=self.device) self.grad_count = torch.zeros(N, device=self.device, dtype=torch.int32) self.max_radii = torch.zeros(N, device=self.device) self._build_optimizer() if "optimizer_state" in state: self.optimizer.load_state_dict(state["optimizer_state"]) print(f"[Checkpoint] Loaded step {self.step} with {N:,} Gaussians") # ------------------------------------------------------------------- # Export # ------------------------------------------------------------------- def export_ply(self, path: str, timestamp: float = 0.0): """ Export deformed Gaussians at a given timestamp to PLY format. Compatible with standard 3DGS viewers. """ from plyfile import PlyElement, PlyData with torch.no_grad(): d_means, d_quats, d_scales, d_opacs = self.apply_deformation(timestamp) means_np = d_means.cpu().numpy() scales_np = torch.exp(d_scales).cpu().numpy() quats_np = F.normalize(d_quats, dim=-1).cpu().numpy() opacities_np = torch.sigmoid(d_opacs).cpu().numpy() sh_np = self.sh_coeffs.data.cpu().numpy() N = means_np.shape[0] num_sh = sh_np.shape[1] # Build structured array dtype = [ ("x", "f4"), ("y", "f4"), ("z", "f4"), ("opacity", "f4"), ("scale_0", "f4"), ("scale_1", "f4"), ("scale_2", "f4"), ("rot_0", "f4"), ("rot_1", "f4"), ("rot_2", "f4"), ("rot_3", "f4"), ] for i in range(num_sh * 3): dtype.append((f"f_rest_{i}", "f4")) arr = np.zeros(N, dtype=dtype) arr["x"] = means_np[:, 0] arr["y"] = means_np[:, 1] arr["z"] = means_np[:, 2] arr["opacity"] = opacities_np[:, 0] arr["scale_0"] = np.log(scales_np[:, 0]) arr["scale_1"] = np.log(scales_np[:, 1]) arr["scale_2"] = np.log(scales_np[:, 2]) arr["rot_0"] = quats_np[:, 0] arr["rot_1"] = quats_np[:, 1] arr["rot_2"] = quats_np[:, 2] arr["rot_3"] = quats_np[:, 3] sh_flat = sh_np.reshape(N, -1) for i in range(sh_flat.shape[1]): arr[f"f_rest_{i}"] = sh_flat[:, i] el = PlyElement.describe(arr, "vertex") PlyData([el]).write(str(path)) print(f"[Export] Saved {N:,} Gaussians to {path}")