From fb8409b75d6f5ba833488d7feb491f62845c6d17 Mon Sep 17 00:00:00 2001 From: Bruno Fargnoli Date: Thu, 5 Mar 2026 16:35:18 +0100 Subject: [PATCH] Added "Heun" sampler --- nodes.py | 12 +-- trellis2/pipelines/samplers/__init__.py | 6 +- trellis2/pipelines/samplers/flow_euler.py | 106 ++++++++++++++++++++- trellis2/pipelines/trellis2_image_to_3d.py | 2 + 4 files changed, 118 insertions(+), 8 deletions(-) diff --git a/nodes.py b/nodes.py index d92d377..7204ffb 100644 --- a/nodes.py +++ b/nodes.py @@ -425,7 +425,7 @@ class Trellis2MeshWithVoxelGenerator: "sparse_structure_resolution": ("INT", {"default":32,"min":32,"max":128,"step":4}), "generate_texture_slat": ("BOOLEAN", {"default":True}), "use_tiled_decoder": ("BOOLEAN", {"default":True}), - "sampler": (["euler", "rk4", "rk5"], {"default": "euler"}), + "sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}), }, } @@ -1234,7 +1234,7 @@ class Trellis2MeshWithVoxelAdvancedGenerator: "texture_guidance_interval_start": ("FLOAT",{"default":0.00,"min":0.00,"max":1.00,"step":0.01}), "texture_guidance_interval_end": ("FLOAT",{"default":0.90,"min":0.00,"max":1.00,"step":0.01}), "use_tiled_decoder": ("BOOLEAN", {"default":True}), - "sampler": (["euler", "rk4", "rk5"], {"default": "euler"}), + "sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}), }, } @@ -1339,7 +1339,7 @@ class Trellis2MeshWithVoxelMultiViewGenerator: "use_tiled_decoder": ("BOOLEAN", {"default":True}), "front_axis": (["z", "x"], {"default": "z"}), "blend_temperature": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), - "sampler": (["euler", "rk4", "rk5"], {"default": "euler"}), + "sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}), }, "optional": { "back_image": ("IMAGE",), @@ -2092,7 +2092,7 @@ class Trellis2MeshTexturing: "bake_on_vertices": ("BOOLEAN",{"default":False}), "use_custom_normals": ("BOOLEAN",{"default":False}), "mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60.0,"min":0.0,"max":359.9}), - "sampler": (["euler", "rk4", "rk5"], {"default": "euler"}), + "sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}), }, } @@ -2156,7 +2156,7 @@ class Trellis2MeshTexturingMultiView: "mesh_cluster_threshold_cone_half_angle_rad": ("FLOAT",{"default":60.0,"min":0.0,"max":359.9}), "front_axis": (["z", "x"], {"default": "z"}), "blend_temperature": ("FLOAT", {"default": 1.0, "min": 0.1, "max": 10.0, "step": 0.1}), - "sampler": (["euler", "rk4", "rk5"], {"default": "euler"}), + "sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}), }, "optional": { "back_image": ("IMAGE",), @@ -2372,7 +2372,7 @@ class Trellis2MeshRefiner: "texture_guidance_interval_end": ("FLOAT",{"default":0.90,"min":0.00,"max":1.00,"step":0.01}), "use_tiled_decoder": ("BOOLEAN", {"default":True}), "max_views": ("INT", {"default": 4, "min": 1, "max": 16}), - "sampler": (["euler", "rk4", "rk5"], {"default": "euler"}), + "sampler": (["euler", "heun", "rk4", "rk5"], {"default": "euler"}), }, } diff --git a/trellis2/pipelines/samplers/__init__.py b/trellis2/pipelines/samplers/__init__.py index 3826607..2d724e4 100644 --- a/trellis2/pipelines/samplers/__init__.py +++ b/trellis2/pipelines/samplers/__init__.py @@ -14,5 +14,9 @@ from .flow_euler import ( FlowRK4MultiViewSampler, FlowRK5MultiViewSampler, FlowRK4MultiViewGuidanceIntervalSampler, - FlowRK5MultiViewGuidanceIntervalSampler + FlowRK5MultiViewGuidanceIntervalSampler, + FlowHeunSampler, + FlowHeunGuidanceIntervalSampler, + FlowHeunMultiViewSampler, + FlowHeunMultiViewGuidanceIntervalSampler ) \ No newline at end of file diff --git a/trellis2/pipelines/samplers/flow_euler.py b/trellis2/pipelines/samplers/flow_euler.py index 73ed91e..fe948aa 100644 --- a/trellis2/pipelines/samplers/flow_euler.py +++ b/trellis2/pipelines/samplers/flow_euler.py @@ -630,4 +630,108 @@ class FlowRK4MultiViewGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, Clas pass class FlowRK5MultiViewGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, ClassifierFreeGuidanceSamplerMixin, FlowRK5MultiViewSampler): - pass \ No newline at end of file + pass + +# Heun (RK2) + +class FlowHeunSampler(FlowEulerSampler): + """ + Generate samples from a flow-matching model using Heun's Method (2nd-order Runge-Kutta). + Requires 2 NFEs per step. + """ + @torch.no_grad() + def sample_once( + self, + model, + x_t, + t: float, + t_prev: float, + cond: Optional[Any] = None, + **kwargs + ): + dt = t_prev - t + + # Helper to extract just the velocity prediction + def get_v(current_x, current_t): + _, _, pred_v = self._get_model_prediction(model, current_x, current_t, cond, **kwargs) + return pred_v + + # Step 1: Predictor (Euler step) + k1 = get_v(x_t, t) + x_temp = x_t + k1 * dt + + # Step 2: Corrector + k2 = get_v(x_temp, t + dt) + + # Average the two velocities for the final step + pred_x_prev = x_t + 0.5 * dt * (k1 + k2) + + # Estimate x_0 based on k1 for tracking/logging + pred_x_0, _ = self._v_to_xstart_eps(x_t=x_t, t=t, v=k1) + + return edict({"pred_x_prev": pred_x_prev, "pred_x_0": pred_x_0}) + +# --- CFG Wrapper for Heun --- +class FlowHeunGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, ClassifierFreeGuidanceSamplerMixin, FlowHeunSampler): + """Heun sampling with CFG and Guidance Intervals.""" + pass + +class FlowHeunMultiViewSampler(FlowEulerMultiViewSampler): + """Multi-view flow matching using Heun's method (2nd-order Runge-Kutta).""" + @torch.no_grad() + def sample_once( + self, model, x_t, t: float, t_prev: float, + conds: Dict[str, Any], views: List[str], + front_axis: str = 'z', blend_temperature: float = 2.0, **kwargs + ): + dt = t_prev - t + is_sparse = hasattr(x_t, 'coords') + + # Calculate spatial blending weights ONCE for the current step + if is_sparse: + weights = self._compute_view_weights_sparse(x_t.coords, views, front_axis, blend_temperature) + else: + weights = self._compute_view_weights_dense(x_t.shape, x_t.device, views, front_axis, blend_temperature) + + # Helper function to compute the blended velocity for a given intermediate x and t + def get_blended_v(current_x, current_t): + pred_v_accum = 0 + for i, view in enumerate(views): + cond = conds[view] + if isinstance(cond, dict) and 'cond' in cond and 'neg_cond' in cond: + pred_v_view = self._inference_model(model, current_x, current_t, cond=cond['cond'], neg_cond=cond['neg_cond'], **kwargs) + else: + pred_v_view = self._inference_model(model, current_x, current_t, cond=cond, **kwargs) + + if is_sparse: + w = weights[:, i].unsqueeze(1) + v_feats = pred_v_view.feats if hasattr(pred_v_view, 'feats') else pred_v_view + pred_v_accum += v_feats * w + else: + w = weights[i].unsqueeze(0).unsqueeze(0) + pred_v_accum += pred_v_view * w + + if is_sparse: + return current_x.replace(feats=pred_v_accum) + else: + return pred_v_accum + + # Heun's Method (RK2) Evaluations + # Step 1: Predictor (Euler step) + k1 = get_blended_v(x_t, t) + x_temp = x_t + k1 * dt + + # Step 2: Corrector + k2 = get_blended_v(x_temp, t + dt) + + # Combine + pred_x_prev = x_t + 0.5 * dt * (k1 + k2) + + # Estimate x_0 based on k1 for tracking + pred_x_0, _ = self._v_to_xstart_eps(x_t=x_t, t=t, v=k1) + + return edict({"pred_x_prev": pred_x_prev, "pred_x_0": pred_x_0}) + +# --- CFG Wrapper for Heun Multi-View --- +class FlowHeunMultiViewGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, ClassifierFreeGuidanceSamplerMixin, FlowHeunMultiViewSampler): + pass \ No newline at end of file diff --git a/trellis2/pipelines/trellis2_image_to_3d.py b/trellis2/pipelines/trellis2_image_to_3d.py index 8951e1a..559b198 100644 --- a/trellis2/pipelines/trellis2_image_to_3d.py +++ b/trellis2/pipelines/trellis2_image_to_3d.py @@ -113,6 +113,8 @@ class Trellis2ImageTo3DPipeline(Pipeline): self._sampler_prefix = "RK4" elif sampler_type == "rk5": self._sampler_prefix = "RK5" + elif sampler_type == "heun": + self._sampler_prefix = "Heun" args = self._pretrained_args self.sparse_structure_sampler = getattr(samplers, f"Flow{self._sampler_prefix}GuidanceIntervalSampler")(**args['sparse_structure_sampler']['args'])