Added "Heun" sampler
This commit is contained in:
@@ -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"}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -14,5 +14,9 @@ from .flow_euler import (
|
||||
FlowRK4MultiViewSampler,
|
||||
FlowRK5MultiViewSampler,
|
||||
FlowRK4MultiViewGuidanceIntervalSampler,
|
||||
FlowRK5MultiViewGuidanceIntervalSampler
|
||||
FlowRK5MultiViewGuidanceIntervalSampler,
|
||||
FlowHeunSampler,
|
||||
FlowHeunGuidanceIntervalSampler,
|
||||
FlowHeunMultiViewSampler,
|
||||
FlowHeunMultiViewGuidanceIntervalSampler
|
||||
)
|
||||
@@ -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'])
|
||||
|
||||
Reference in New Issue
Block a user