Added "Heun" sampler

This commit is contained in:
Bruno Fargnoli
2026-03-05 16:35:18 +01:00
parent 4b6cfc43e4
commit fb8409b75d
4 changed files with 118 additions and 8 deletions
+6 -6
View File
@@ -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"}),
},
}
+5 -1
View File
@@ -14,5 +14,9 @@ from .flow_euler import (
FlowRK4MultiViewSampler,
FlowRK5MultiViewSampler,
FlowRK4MultiViewGuidanceIntervalSampler,
FlowRK5MultiViewGuidanceIntervalSampler
FlowRK5MultiViewGuidanceIntervalSampler,
FlowHeunSampler,
FlowHeunGuidanceIntervalSampler,
FlowHeunMultiViewSampler,
FlowHeunMultiViewGuidanceIntervalSampler
)
+105 -1
View File
@@ -630,4 +630,108 @@ class FlowRK4MultiViewGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, Clas
pass
class FlowRK5MultiViewGuidanceIntervalSampler(GuidanceIntervalSamplerMixin, ClassifierFreeGuidanceSamplerMixin, FlowRK5MultiViewSampler):
pass
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
@@ -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'])