Spaces:
Running on Zero
Running on Zero
| """Interactive 3DGS + MPM simulation session (extracted from simulation_gt.py).""" | |
| import json | |
| import math | |
| import os | |
| import sys | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import warp as wp | |
| sys.path.append("gs") | |
| from scene.gaussian_model import GaussianModel | |
| sys.path.append("utils") | |
| from utils.decode_param import decode_param_json, find_far_points, set_boundary_conditions | |
| from utils.gpu_runtime import ensure_taichi_runtime, release_taichi_runtime | |
| def _hf_render_max_gaussians() -> int: | |
| return int(os.environ.get("ENDOGSIM_HF_MAX_GAUSSIANS", "150000")) | |
| from utils.transformation_utils import * | |
| from utils.camera_view_utils import * | |
| from utils.render_utils import * | |
| from mpm_solver_warp.engine_utils import * | |
| from mpm_solver_warp.mpm_solver_warp import MPM_Simulator_WARP | |
| from particle_filling.filling import * | |
| class PipelineParamsNoparse: | |
| def __init__(self): | |
| self.convert_SHs_python = False | |
| self.compute_cov3D_python = False | |
| self.debug = False | |
| def load_checkpoint(model_path, iteration=-1, material=None, ply_name="point_cloud.ply", dataset="endonerf"): | |
| checkpt_dir = os.path.join(model_path, "point_cloud") | |
| if iteration == -1: | |
| iteration = searchForMaxIteration(checkpt_dir) | |
| if dataset in ("endonerf", "cholecseg_sub", "porcine_endo"): | |
| checkpt_path = os.path.join(checkpt_dir, f"iteration_{iteration}", ply_name) | |
| else: | |
| checkpt_path = os.path.join(checkpt_dir, f"iteration_{iteration}", "point_cloud.ply") | |
| from plyfile import PlyData | |
| plydata = PlyData.read(checkpt_path) | |
| extra_f_names = [p.name for p in plydata.elements[0].properties if p.name.startswith("f_rest_")] | |
| extra_f_names = sorted(extra_f_names, key=lambda x: int(x.split("_")[-1])) | |
| sh_degree = int(math.sqrt((len(extra_f_names) + 3) // 3)) - 1 | |
| gaussians = GaussianModel(sh_degree) | |
| gaussians.load_ply(checkpt_path, material) | |
| return gaussians | |
| def load_inpaint_gs(model_path): | |
| checkpt_path = os.path.join(model_path, "inpaint_points.ply") | |
| if not os.path.exists(checkpt_path): | |
| return None | |
| from plyfile import PlyData | |
| plydata = PlyData.read(checkpt_path) | |
| extra_f_names = [p.name for p in plydata.elements[0].properties if p.name.startswith("f_rest_")] | |
| extra_f_names = sorted(extra_f_names, key=lambda x: int(x.split("_")[-1])) | |
| sh_degree = int(math.sqrt((len(extra_f_names) + 3) // 3)) - 1 | |
| gaussians = GaussianModel(sh_degree) | |
| gaussians.load_ply(checkpt_path) | |
| return gaussians | |
| FIXED_CAMERA_DATASETS = ("endonerf", "cholecseg_sub", "porcine_endo") | |
| ORBIT_AZIMUTH_MIN = -180.0 | |
| ORBIT_AZIMUTH_MAX = 180.0 | |
| ORBIT_ELEVATION_MIN = -89.0 | |
| ORBIT_ELEVATION_MAX = 89.0 | |
| ORBIT_RADIUS_MIN = 0.1 | |
| ORBIT_RADIUS_MAX = 100.0 | |
| def _sanitize_orbit_value(value, default, lo, hi): | |
| v = float(value) | |
| if not np.isfinite(v): | |
| v = float(default) | |
| return float(np.clip(v, lo, hi)) | |
| def _get_dataset_fixed_camera(dataset): | |
| if dataset == "endonerf": | |
| return get_camera_view_endonerf() | |
| if dataset == "cholecseg_sub": | |
| return get_camera_view_cholecseg_sub() | |
| if dataset == "porcine_endo": | |
| return get_camera_view_porcine_endo() | |
| raise ValueError(f"No fixed camera for dataset: {dataset}") | |
| def _intrinsics_from_fixed_camera(camera, downsample=1.0): | |
| width = max(1, int(camera.image_width * downsample)) | |
| height = max(1, int(camera.image_height * downsample)) | |
| return width, height, camera.FoVx, camera.FoVy | |
| def _clone_fixed_camera_with_downsample(fixed_cam, downsample=1.0): | |
| """Reuse dataset fixed extrinsics (same as simulation_gt.py); only scale resolution.""" | |
| from scene.cameras import Camera as GSCamera | |
| width, height, fovx, fovy = _intrinsics_from_fixed_camera(fixed_cam, downsample) | |
| return GSCamera( | |
| colmap_id=fixed_cam.colmap_id, | |
| R=np.array(fixed_cam.R, copy=True), | |
| T=np.array(fixed_cam.T, copy=True), | |
| FoVx=fovx, | |
| FoVy=fovy, | |
| image_width=width, | |
| image_height=height, | |
| image=torch.zeros((3, height, width)), | |
| gt_alpha_mask=None, | |
| image_name=fixed_cam.image_name, | |
| image_path=fixed_cam.image_path, | |
| uid=fixed_cam.uid, | |
| preload_img=False, | |
| ) | |
| def _orbit_from_camera(camera, viewpoint_center, observant_coordinates): | |
| cam_pos = camera.camera_center.detach().cpu().numpy() | |
| radius, azimuth, elevation = get_current_radius_azimuth_and_elevation( | |
| cam_pos, viewpoint_center, observant_coordinates | |
| ) | |
| return azimuth, elevation, radius | |
| def _load_camera_intrinsics(model_path, default_camera_index=0, downsample=1.0): | |
| cam_path = os.path.join(model_path, "cameras.json") | |
| with open(cam_path) as f: | |
| data = json.load(f) | |
| raw = data[default_camera_index] if default_camera_index > -1 else data[0] | |
| width = int(min(raw["width"], 1920) * downsample) | |
| height = int(min(raw["height"], 1920) * downsample) | |
| from utils.graphics_utils import focal2fov | |
| fovx = focal2fov(raw["fx"] * downsample, width) | |
| fovy = focal2fov(raw["fy"] * downsample, height) | |
| return width, height, fovx, fovy | |
| def build_orbit_camera( | |
| width, | |
| height, | |
| fovx, | |
| fovy, | |
| azimuth, | |
| elevation, | |
| radius, | |
| viewpoint_center, | |
| observant_coordinates, | |
| ): | |
| from scene.cameras import Camera as GSCamera | |
| position, R = get_camera_position_and_rotation( | |
| azimuth, elevation, radius, viewpoint_center, observant_coordinates | |
| ) | |
| tmp = np.zeros((4, 4)) | |
| tmp[:3, :3] = R | |
| tmp[:3, 3] = position | |
| tmp[3, 3] = 1 | |
| c2w = np.linalg.inv(tmp) | |
| cam_R = c2w[:3, :3].transpose() | |
| cam_T = c2w[:3, 3] | |
| return GSCamera( | |
| colmap_id=0, | |
| R=cam_R, | |
| T=cam_T, | |
| FoVx=fovx, | |
| FoVy=fovy, | |
| image_width=width, | |
| image_height=height, | |
| image=torch.zeros((3, height, width)), | |
| gt_alpha_mask=None, | |
| image_name="interactive", | |
| image_path="interactive", | |
| uid=0, | |
| preload_img=False, | |
| ) | |
| class SimulationSession: | |
| """Holds MPM state, 3DGS assets, and camera for interactive stepping.""" | |
| def __init__( | |
| self, | |
| model_path, | |
| physics_config, | |
| dataset="pacnerf", | |
| white_bg=False, | |
| downsample=0.5, | |
| ply_name="point_cloud.ply", | |
| ): | |
| self.model_path = model_path | |
| self.dataset = dataset | |
| self.device = "cuda:0" | |
| self.downsample = downsample | |
| self.white_bg = white_bg | |
| ( | |
| material_params, | |
| bc_params, | |
| time_params, | |
| preprocessing_params, | |
| camera_params, | |
| _optimize_params, | |
| ) = decode_param_json(physics_config) | |
| self.material_params = material_params | |
| self.bc_params = bc_params | |
| self.time_params = time_params | |
| self.preprocessing_params = preprocessing_params | |
| self.camera_params = camera_params | |
| self.substep_dt = time_params["substep_dt"] | |
| self.substeps_per_frame = max(1, int(time_params["frame_dt"] / time_params["substep_dt"])) | |
| self._hf_gsplat_logged = False | |
| gaussians = load_checkpoint( | |
| model_path, material=material_params["material"], ply_name=ply_name, dataset=dataset | |
| ) | |
| gaussians_inpaint = load_inpaint_gs(model_path) | |
| pipeline = PipelineParamsNoparse() | |
| pipeline.compute_cov3D_python = True | |
| self.pipeline = pipeline | |
| self.gaussians = gaussians | |
| self.background = ( | |
| torch.tensor([1, 1, 1], dtype=torch.float32, device="cuda") | |
| if white_bg | |
| else torch.tensor([0, 0, 0], dtype=torch.float32, device="cuda") | |
| ) | |
| params = load_params_from_gs(gaussians, pipeline) | |
| params_inpaint = load_params_from_gs(gaussians_inpaint, pipeline) if gaussians_inpaint else None | |
| self.params_inpaint = params_inpaint | |
| init_pos = params["pos"] | |
| init_cov = params["cov3D_precomp"] | |
| init_screen_points = params["screen_points"] | |
| init_opacity = params["opacity"] | |
| init_shs = params["shs"] | |
| mask = init_opacity[:, 0] > preprocessing_params["opacity_threshold"] | |
| init_pos = init_pos[mask] | |
| init_cov = init_cov[mask] | |
| init_opacity = init_opacity[mask] | |
| init_screen_points = init_screen_points[mask] | |
| init_shs = init_shs[mask] | |
| unselected_pos = unselected_cov = unselected_opacity = unselected_shs = None | |
| moving_pts_path = os.path.join(model_path, "moving_part_points.ply") | |
| self.moving_pts_path = moving_pts_path | |
| if os.path.exists(moving_pts_path): | |
| import point_cloud_utils as pcu | |
| moving_pts = torch.from_numpy(pcu.load_mesh_v(moving_pts_path)).float().to("cuda") | |
| thres = 0.5 / material_params["n_grid"] | |
| if "playdoh" in model_path: | |
| thres = 1.0 / material_params["n_grid"] | |
| freeze_mask = find_far_points(init_pos, moving_pts, thres=thres).bool() | |
| unselected_pos = init_pos[freeze_mask] | |
| unselected_cov = init_cov[freeze_mask] | |
| unselected_opacity = init_opacity[freeze_mask] | |
| unselected_shs = init_shs[freeze_mask] | |
| init_pos = init_pos[~freeze_mask] | |
| init_cov = init_cov[~freeze_mask] | |
| init_opacity = init_opacity[~freeze_mask] | |
| init_shs = init_shs[~freeze_mask] | |
| rotation_matrices = generate_rotation_matrices( | |
| torch.tensor(preprocessing_params["rotation_degree"]), | |
| preprocessing_params["rotation_axis"], | |
| ) | |
| self.rotation_matrices = rotation_matrices | |
| rotated_pos = apply_rotations(init_pos, rotation_matrices) | |
| if preprocessing_params["sim_area"] is not None: | |
| boundary = preprocessing_params["sim_area"] | |
| area_mask = torch.ones(rotated_pos.shape[0], dtype=torch.bool, device="cuda") | |
| for i in range(3): | |
| area_mask = torch.logical_and(area_mask, rotated_pos[:, i] > boundary[2 * i]) | |
| area_mask = torch.logical_and(area_mask, rotated_pos[:, i] < boundary[2 * i + 1]) | |
| unselected_pos = init_pos[~area_mask] | |
| unselected_cov = init_cov[~area_mask] | |
| unselected_opacity = init_opacity[~area_mask] | |
| unselected_shs = init_shs[~area_mask] | |
| rotated_pos = rotated_pos[area_mask] | |
| init_cov = init_cov[area_mask] | |
| init_opacity = init_opacity[area_mask] | |
| init_shs = init_shs[area_mask] | |
| scaling = 1.0 | |
| for key, val in [("cat", 0.7), ("letter", 2.0), ("cream", 0.8), ("toothpaste", 0.6), ("playdoh", 0.75)]: | |
| if key in model_path: | |
| scaling = val | |
| transformed_pos, scale_origin, original_mean_pos = transform2origin(rotated_pos, scaling=scaling) | |
| transformed_pos = shift2center111(transformed_pos) | |
| self.scale_origin = scale_origin | |
| self.original_mean_pos = original_mean_pos | |
| init_cov = apply_cov_rotations(init_cov, rotation_matrices) | |
| init_cov = scale_origin * scale_origin * init_cov | |
| gs_num = transformed_pos.shape[0] | |
| ensure_taichi_runtime() | |
| filling_params = preprocessing_params["particle_filling"] | |
| if filling_params is not None: | |
| mpm_init_pos = fill_particles( | |
| pos=transformed_pos, | |
| opacity=init_opacity, | |
| cov=init_cov, | |
| grid_n=filling_params["n_grid"], | |
| max_samples=filling_params["max_particles_num"], | |
| grid_dx=material_params["grid_lim"] / filling_params["n_grid"], | |
| density_thres=filling_params["density_threshold"], | |
| search_thres=filling_params["search_threshold"], | |
| max_particles_per_cell=filling_params["max_partciels_per_cell"], | |
| search_exclude_dir=filling_params["search_exclude_direction"], | |
| ray_cast_dir=filling_params["ray_cast_direction"], | |
| boundary=filling_params["boundary"], | |
| smooth=filling_params["smooth"], | |
| ).to(device=self.device) | |
| else: | |
| mpm_init_pos = transformed_pos.to(device=self.device) | |
| mpm_init_vol = get_particle_volume( | |
| mpm_init_pos, | |
| material_params["n_grid"], | |
| material_params["grid_lim"] / material_params["n_grid"], | |
| unifrom=material_params["material"] == "sand", | |
| ).to(device=self.device) | |
| if filling_params is not None and filling_params.get("visualize", False): | |
| shs, opacity, mpm_init_cov = init_filled_particles( | |
| mpm_init_pos[:gs_num], init_shs, init_cov, init_opacity, mpm_init_pos[gs_num:] | |
| ) | |
| _pos = apply_inverse_rotations( | |
| undotransform2origin( | |
| undoshift2center111(mpm_init_pos[gs_num:]), scale_origin, original_mean_pos | |
| ), | |
| rotation_matrices, | |
| ) | |
| gaussians._xyz = nn.Parameter( | |
| torch.cat([gaussians._xyz, _pos], 0).float().cuda().requires_grad_(True) | |
| ) | |
| gaussians._opacity = nn.Parameter( | |
| torch.cat([gaussians._opacity, torch.zeros((_pos.shape[0], 1), device="cuda")], 0) | |
| .float() | |
| .cuda() | |
| .requires_grad_(True) | |
| ) | |
| gaussians._scaling = nn.Parameter( | |
| torch.cat([gaussians._scaling, torch.zeros((_pos.shape[0], 3), device="cuda")], 0) | |
| .float() | |
| .cuda() | |
| .requires_grad_(True) | |
| ) | |
| gaussians._rotation = nn.Parameter( | |
| torch.cat([gaussians._rotation, torch.zeros((_pos.shape[0], 4), device="cuda")], 0) | |
| .float() | |
| .cuda() | |
| .requires_grad_(True) | |
| ) | |
| gs_num = mpm_init_pos.shape[0] | |
| else: | |
| mpm_init_cov = torch.zeros((mpm_init_pos.shape[0], 6), device=self.device) | |
| mpm_init_cov[:gs_num] = init_cov | |
| shs = init_shs | |
| opacity = init_opacity | |
| self.gs_num = gs_num | |
| self.init_len = mpm_init_pos.shape[0] | |
| self.init_screen_points = init_screen_points | |
| self.opacity_render = opacity | |
| self.shs_render = shs | |
| self.unselected_pos = unselected_pos | |
| self.unselected_cov = unselected_cov | |
| self.unselected_opacity = unselected_opacity | |
| self.unselected_shs = unselected_shs | |
| mpm_solver = MPM_Simulator_WARP(10) | |
| mpm_solver.load_initial_data_from_torch( | |
| mpm_init_pos, | |
| mpm_init_vol, | |
| mpm_init_cov, | |
| n_grid=material_params["n_grid"], | |
| grid_lim=material_params["grid_lim"], | |
| ) | |
| mpm_solver.set_parameters_dict(material_params) | |
| if dataset in ("endonerf", "cholecseg_sub", "porcine_endo"): | |
| for bc in bc_params: | |
| if bc["type"] in ("particle_impulse", "cuboid"): | |
| bc["point"] = bc["point"] - original_mean_pos.detach().cpu().numpy() | |
| bc["point"] = bc["point"] * scale_origin.detach().cpu().numpy() | |
| bc["point"] = bc["point"] + np.array([1.0, 1.0, 1.0]) | |
| if bc["type"] in ("particle_velocity",): | |
| bc["size"] = bc["size"] * scale_origin.detach().cpu().numpy() | |
| set_boundary_conditions(mpm_solver, bc_params, time_params) | |
| mpm_solver.finalize_mu_lam() | |
| self.mpm_solver = mpm_solver | |
| self._save_initial_state(mpm_init_pos, mpm_init_cov, mpm_init_vol) | |
| mpm_space_viewpoint_center = ( | |
| torch.tensor(camera_params["mpm_space_viewpoint_center"]).reshape((1, 3)).cuda() | |
| ) | |
| mpm_space_vertical_upward_axis = ( | |
| torch.tensor(camera_params["mpm_space_vertical_upward_axis"]).reshape((1, 3)).cuda() | |
| ) | |
| viewpoint_center, observant_coordinates = get_center_view_worldspace_and_observant_coordinate( | |
| mpm_space_viewpoint_center, | |
| mpm_space_vertical_upward_axis, | |
| rotation_matrices, | |
| scale_origin, | |
| original_mean_pos, | |
| ) | |
| self.viewpoint_center = viewpoint_center | |
| self.observant_coordinates = observant_coordinates | |
| if dataset in FIXED_CAMERA_DATASETS: | |
| # Match simulation_gt.py: use dataset fixed camera directly (no orbit rebuild). | |
| self.use_fixed_camera = True | |
| fixed_cam = _get_dataset_fixed_camera(dataset) | |
| self.cam_width, self.cam_height, self.cam_fovx, self.cam_fovy = ( | |
| _intrinsics_from_fixed_camera(fixed_cam, downsample) | |
| ) | |
| self.current_camera = _clone_fixed_camera_with_downsample(fixed_cam, downsample) | |
| self.rasterize = initialize_resterize( | |
| self.current_camera, self.gaussians, self.pipeline, self.background | |
| ) | |
| else: | |
| self.use_fixed_camera = False | |
| self.cam_width, self.cam_height, self.cam_fovx, self.cam_fovy = _load_camera_intrinsics( | |
| model_path, camera_params["default_camera_index"], downsample | |
| ) | |
| _, cam_info = get_camera_view( | |
| model_path, | |
| default_camera_index=camera_params["default_camera_index"], | |
| center_view_world_space=viewpoint_center, | |
| observant_coordinates=observant_coordinates, | |
| downsample=downsample, | |
| ) | |
| self.orbit_azimuth = cam_info["init_azimuthm"] | |
| self.orbit_elevation = cam_info["init_elevation"] | |
| self.orbit_radius = cam_info["init_radius"] | |
| self.set_orbit( | |
| azimuth=self.orbit_azimuth, | |
| elevation=self.orbit_elevation, | |
| radius=self.orbit_radius, | |
| ) | |
| release_taichi_runtime() | |
| self.frame_idx = 0 | |
| def _save_initial_state(self, mpm_init_pos, mpm_init_cov, mpm_init_vol): | |
| self._init_mpm_pos = mpm_init_pos.clone() | |
| self._init_mpm_cov = mpm_init_cov.clone() | |
| self._init_mpm_vol = mpm_init_vol.clone() | |
| def reset_simulation(self): | |
| # Reuse cached volume; avoid Taichi get_particle_volume in Gradio worker threads. | |
| self.mpm_solver.reset_pos_from_torch( | |
| self._init_mpm_pos, | |
| self._init_mpm_vol, | |
| self._init_mpm_cov, | |
| device=self.device, | |
| ) | |
| self.frame_idx = 0 | |
| def _update_camera(self): | |
| if getattr(self, "use_fixed_camera", False): | |
| fixed_cam = _get_dataset_fixed_camera(self.dataset) | |
| self.current_camera = _clone_fixed_camera_with_downsample( | |
| fixed_cam, self.downsample | |
| ) | |
| else: | |
| self.current_camera = build_orbit_camera( | |
| self.cam_width, | |
| self.cam_height, | |
| self.cam_fovx, | |
| self.cam_fovy, | |
| self.orbit_azimuth, | |
| self.orbit_elevation, | |
| self.orbit_radius, | |
| self.viewpoint_center, | |
| self.observant_coordinates, | |
| ) | |
| self.rasterize = initialize_resterize( | |
| self.current_camera, self.gaussians, self.pipeline, self.background | |
| ) | |
| def set_orbit(self, azimuth=None, elevation=None, radius=None): | |
| if azimuth is not None: | |
| self.orbit_azimuth = _sanitize_orbit_value( | |
| azimuth, 0.0, ORBIT_AZIMUTH_MIN, ORBIT_AZIMUTH_MAX | |
| ) | |
| if elevation is not None: | |
| self.orbit_elevation = _sanitize_orbit_value( | |
| elevation, 15.0, ORBIT_ELEVATION_MIN, ORBIT_ELEVATION_MAX | |
| ) | |
| if radius is not None: | |
| self.orbit_radius = _sanitize_orbit_value( | |
| radius, 2.0, ORBIT_RADIUS_MIN, ORBIT_RADIUS_MAX | |
| ) | |
| self._update_camera() | |
| def step(self, n_substeps=None): | |
| n = n_substeps if n_substeps is not None else self.substeps_per_frame | |
| for _ in range(n): | |
| self.mpm_solver.p2g2p(self.frame_idx, self.substep_dt, device=self.device) | |
| self.frame_idx += 1 | |
| def get_world_positions(self): | |
| pos = self.mpm_solver.export_particle_x_to_torch()[: self.gs_num].to(self.device) | |
| pos = pos[: self.init_len] | |
| return undo_all_transforms(pos, self.rotation_matrices, self.scale_origin, self.original_mean_pos) | |
| def max_particle_speed(self) -> float: | |
| """Max |v| over sim particles (MPM space); used for HF burst early-stop.""" | |
| v = self.mpm_solver.export_particle_v_to_torch()[: self.init_len] | |
| if v is None or v.numel() == 0: | |
| return 0.0 | |
| return float(torch.linalg.vector_norm(v, dim=-1).max().item()) | |
| def apply_impulse_at_world(self, world_point, force, radius=0.05, num_dt=20): | |
| world = torch.tensor(world_point, dtype=torch.float32, device="cuda").reshape(1, 3) | |
| mpm_point = world_to_mpm(world, self.rotation_matrices, self.scale_origin, self.original_mean_pos) | |
| point = mpm_point.detach().cpu().numpy().reshape(-1).tolist() | |
| size = [radius, radius, radius] | |
| self.mpm_solver.add_impulse_on_particles( | |
| force=force, | |
| dt=self.substep_dt, | |
| point=point, | |
| size=size, | |
| num_dt=num_dt, | |
| start_time=self.mpm_solver.time, | |
| device=self.device, | |
| ) | |
| def _hf_static_gaussian_count(self) -> int: | |
| """How many static (non-sim) gaussians _append_static_render_gaussians will add.""" | |
| n = 0 | |
| if os.path.exists(self.moving_pts_path) and self.unselected_pos is not None: | |
| n += int(self.unselected_pos.shape[0]) | |
| if self.params_inpaint is not None: | |
| n += int(self.params_inpaint["pos"].shape[0]) | |
| if self.preprocessing_params["sim_area"] is not None and self.unselected_pos is not None: | |
| n += int(self.unselected_pos.shape[0]) | |
| return n | |
| def _append_static_render_gaussians( | |
| self, | |
| pos: torch.Tensor, | |
| cov3D: torch.Tensor, | |
| opacity: torch.Tensor, | |
| shs: torch.Tensor, | |
| ): | |
| """Concatenate frozen / inpaint background gaussians (same as local full render).""" | |
| if os.path.exists(self.moving_pts_path) and self.unselected_pos is not None: | |
| pos = torch.cat([pos, self.unselected_pos], dim=0) | |
| cov3D = torch.cat([cov3D, self.unselected_cov], dim=0) | |
| opacity = torch.cat([opacity, self.unselected_opacity], dim=0) | |
| shs = torch.cat([shs, self.unselected_shs], dim=0) | |
| if self.params_inpaint is not None: | |
| pos = torch.cat([pos, self.params_inpaint["pos"]], dim=0) | |
| cov3D = torch.cat([cov3D, self.params_inpaint["cov3D_precomp"]], dim=0) | |
| opacity = torch.cat([opacity, self.params_inpaint["opacity"]], dim=0) | |
| shs = torch.cat([shs, self.params_inpaint["shs"]], dim=0) | |
| if self.preprocessing_params["sim_area"] is not None and self.unselected_pos is not None: | |
| pos = torch.cat([pos, self.unselected_pos], dim=0) | |
| cov3D = torch.cat([cov3D, self.unselected_cov], dim=0) | |
| opacity = torch.cat([opacity, self.unselected_opacity], dim=0) | |
| shs = torch.cat([shs, self.unselected_shs], dim=0) | |
| return pos, cov3D, opacity, shs | |
| def _subsample_tensor_rows(self, tensor: torch.Tensor, target: int) -> torch.Tensor: | |
| n = tensor.shape[0] | |
| if target <= 0 or n <= target: | |
| return tensor | |
| idx = torch.linspace(0, n - 1, target, device=tensor.device).long() | |
| return tensor[idx] | |
| def _subsample_hf_sim_tensors( | |
| self, | |
| pos: torch.Tensor, | |
| cov3D: torch.Tensor, | |
| rot: torch.Tensor, | |
| opacity: torch.Tensor, | |
| shs: torch.Tensor, | |
| target: int, | |
| ): | |
| if target <= 0 or pos.shape[0] <= target: | |
| return pos, cov3D, rot, opacity, shs | |
| if not getattr(self, "_hf_subsample_logged", False): | |
| self._hf_subsample_logged = True | |
| print( | |
| f"HF gsplat: subsampling sim particles {pos.shape[0]} -> {target} " | |
| f"(static background kept)" | |
| ) | |
| idx = torch.linspace(0, pos.shape[0] - 1, target, device=pos.device).long() | |
| return pos[idx], cov3D[idx], rot[idx], opacity[idx], shs[idx] | |
| def _cap_hf_static_tail( | |
| self, | |
| pos: torch.Tensor, | |
| cov3D: torch.Tensor, | |
| opacity: torch.Tensor, | |
| shs: torch.Tensor, | |
| n_sim: int, | |
| max_total: int, | |
| ): | |
| """If sim+static exceeds cap, subsample static tail only.""" | |
| if max_total <= 0 or pos.shape[0] <= max_total: | |
| return pos, cov3D, opacity, shs | |
| n_sim = min(n_sim, pos.shape[0]) | |
| static_budget = max(0, max_total - n_sim) | |
| if static_budget <= 0: | |
| return pos[:n_sim], cov3D[:n_sim], opacity[:n_sim], shs[:n_sim] | |
| static_pos = pos[n_sim:] | |
| if static_pos.shape[0] <= static_budget: | |
| return pos, cov3D, opacity, shs | |
| static_pos = self._subsample_tensor_rows(static_pos, static_budget) | |
| static_cov = self._subsample_tensor_rows(cov3D[n_sim:], static_budget) | |
| static_op = self._subsample_tensor_rows(opacity[n_sim:], static_budget) | |
| static_shs = self._subsample_tensor_rows(shs[n_sim:], static_budget) | |
| return ( | |
| torch.cat([pos[:n_sim], static_pos], dim=0), | |
| torch.cat([cov3D[:n_sim], static_cov], dim=0), | |
| torch.cat([opacity[:n_sim], static_op], dim=0), | |
| torch.cat([shs[:n_sim], static_shs], dim=0), | |
| ) | |
| def render(self): | |
| mpm_solver = self.mpm_solver | |
| gs_num = self.gs_num | |
| init_len = self.init_len | |
| hf_lite = os.environ.get("ENDOGSIM_HF_SPACE") == "1" | |
| pos = mpm_solver.export_particle_x_to_torch()[:gs_num].to(self.device) | |
| cov3D = mpm_solver.export_particle_cov_to_torch() | |
| rot = mpm_solver.export_particle_R_to_torch() | |
| cov3D = cov3D.view(-1, 6)[:gs_num].to(self.device) | |
| rot = rot.view(-1, 3, 3)[:gs_num].to(self.device) | |
| pos = pos[:init_len] | |
| cov3D = cov3D[:init_len] | |
| rot = rot[:init_len] | |
| pos = undo_all_transforms(pos, self.rotation_matrices, self.scale_origin, self.original_mean_pos) | |
| cov3D = cov3D / (self.scale_origin * self.scale_origin) | |
| cov3D = apply_inverse_cov_rotations(cov3D, self.rotation_matrices) | |
| n_pts = pos.shape[0] | |
| opacity = self.opacity_render[:n_pts] | |
| shs = self.shs_render[:n_pts] | |
| if hf_lite: | |
| max_g = _hf_render_max_gaussians() | |
| bg_n = self._hf_static_gaussian_count() | |
| if max_g > 0 and n_pts + bg_n > max_g: | |
| sim_target = max(1, max_g - bg_n) | |
| if sim_target < n_pts: | |
| pos, cov3D, rot, opacity, shs = self._subsample_hf_sim_tensors( | |
| pos, cov3D, rot, opacity, shs, sim_target | |
| ) | |
| n_sim = pos.shape[0] | |
| pos, cov3D, opacity, shs = self._append_static_render_gaussians( | |
| pos, cov3D, opacity, shs | |
| ) | |
| if hf_lite: | |
| max_g = _hf_render_max_gaussians() | |
| pos, cov3D, opacity, shs = self._cap_hf_static_tail( | |
| pos, cov3D, opacity, shs, n_sim, max_g | |
| ) | |
| colors_precomp = convert_SH(shs, self.current_camera, self.gaussians, pos, rot) | |
| if hf_lite: | |
| return self._render_with_gsplat(pos, cov3D, opacity, colors_precomp) | |
| if self.init_screen_points.shape[0] == pos.shape[0]: | |
| means2D = self.init_screen_points | |
| else: | |
| means2D = torch.zeros( | |
| (pos.shape[0], *self.init_screen_points.shape[1:]), | |
| dtype=self.init_screen_points.dtype, | |
| device=pos.device, | |
| ) | |
| rendering, _, _, _ = self.rasterize( | |
| means3D=pos, | |
| means2D=means2D, | |
| means2D_abs=pos, | |
| shs=None, | |
| colors_precomp=colors_precomp, | |
| opacities=opacity, | |
| scales=None, | |
| rotations=None, | |
| cov3D_precomp=cov3D, | |
| ) | |
| return rendering | |
| def _cov6_to_covars(self, cov6: torch.Tensor) -> torch.Tensor: | |
| """Upper-triangular cov3D [N,6] -> symmetric [N,3,3] for gsplat.""" | |
| n = cov6.shape[0] | |
| covars = torch.zeros((n, 3, 3), dtype=cov6.dtype, device=cov6.device) | |
| covars[:, 0, 0] = cov6[:, 0] | |
| covars[:, 0, 1] = covars[:, 1, 0] = cov6[:, 1] | |
| covars[:, 0, 2] = covars[:, 2, 0] = cov6[:, 2] | |
| covars[:, 1, 1] = cov6[:, 3] | |
| covars[:, 1, 2] = covars[:, 2, 1] = cov6[:, 4] | |
| covars[:, 2, 2] = cov6[:, 5] | |
| return covars | |
| def _render_with_gsplat( | |
| self, | |
| pos: torch.Tensor, | |
| cov3D: torch.Tensor, | |
| opacity: torch.Tensor, | |
| colors: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """HF ZeroGPU path: MPM unchanged, render via gsplat (avoids broken plane-ext on MIG).""" | |
| from gsplat import rasterization | |
| cam = self.current_camera | |
| height = int(cam.image_height) | |
| width = int(cam.image_width) | |
| # world_view_transform is stored transposed for GLM; gsplat wants row-major w2c. | |
| viewmat = cam.world_view_transform.transpose(0, 1).contiguous().unsqueeze(0) | |
| K = torch.tensor( | |
| [ | |
| [float(cam.Fx), 0.0, float(cam.Cx)], | |
| [0.0, float(cam.Fy), float(cam.Cy)], | |
| [0.0, 0.0, 1.0], | |
| ], | |
| dtype=pos.dtype, | |
| device=pos.device, | |
| ).unsqueeze(0) | |
| opacities = opacity.reshape(-1).float() | |
| colors_rgb = colors.reshape(-1, 3).float().clamp(0.0, 1.0) | |
| covars = self._cov6_to_covars(cov3D.float()) | |
| # Dummy quats/scales ignored when covars is provided. | |
| quats = torch.zeros((pos.shape[0], 4), dtype=pos.dtype, device=pos.device) | |
| quats[:, 0] = 1.0 | |
| scales = torch.ones((pos.shape[0], 3), dtype=pos.dtype, device=pos.device) | |
| bg = self.background | |
| backgrounds = bg.reshape(-1).float() # packed mode expects (channels,), not (1, C) | |
| if not self._hf_gsplat_logged: | |
| self._hf_gsplat_logged = True | |
| print( | |
| f"HF gsplat render: N={pos.shape[0]} hw=({height},{width}) " | |
| f"device={torch.cuda.get_device_name(0)}" | |
| ) | |
| render_colors, _render_alphas, _meta = rasterization( | |
| means=pos.float(), | |
| quats=quats, | |
| scales=scales, | |
| opacities=opacities, | |
| colors=colors_rgb, | |
| viewmats=viewmat.float(), | |
| Ks=K, | |
| width=width, | |
| height=height, | |
| backgrounds=backgrounds, | |
| covars=covars, | |
| packed=True, | |
| render_mode="RGB", | |
| ) | |
| # [C, H, W, 3] or [H, W, 3] depending on version | |
| if render_colors.ndim == 4: | |
| rgb = render_colors[0] | |
| else: | |
| rgb = render_colors | |
| out = rgb.permute(2, 0, 1).contiguous() | |
| return out |