Fix cold start overlap trimming
Should not drop frames on first iter
This commit is contained in:
@@ -107,7 +107,9 @@ class TensorLoopOpen(io.ComfyNode):
|
||||
comfy.utils.ProgressBar(total_frames_val, node_id=open_node_id).update_absolute(accumulated_count)
|
||||
elif count > 0:
|
||||
comfy.utils.ProgressBar(count, node_id=open_node_id).update_absolute(count - remaining)
|
||||
loop_state = {"remaining": remaining, "accum": accum, "previous_value": previous_value, "count": count, "open_node_id": open_node_id, "total_frames": total_frames_val, "blend_overlap": blend_overlap}
|
||||
# The initial_value link survives loop-body cloning, so this is valid on every iteration
|
||||
has_initial = initial_value is not None
|
||||
loop_state = {"remaining": remaining, "accum": accum, "previous_value": previous_value, "count": count, "open_node_id": open_node_id, "total_frames": total_frames_val, "blend_overlap": blend_overlap, "has_initial": has_initial}
|
||||
return io.NodeOutput(loop_state, previous_value, accumulated_count, current_iteration)
|
||||
|
||||
|
||||
@@ -144,10 +146,10 @@ class TensorLoopClose(io.ComfyNode):
|
||||
"- start/end: cut the duplicates from the start or end of each iteration's output\n"
|
||||
"- fade_linear/fade_smooth: crossfade the end of each iteration into the start of the next", options=[
|
||||
io.DynamicCombo.Option("disabled", []),
|
||||
_overlap_option("start", "Number of frames to cut from the START of each iteration's output. Use when the model re-generates its context frames at the start of its output (most common for video continuation)."),
|
||||
_overlap_option("end", "Number of frames to cut from the END of each iteration's output. Use when the model generates look-ahead frames at the end of its output that the next iteration re-generates."),
|
||||
_overlap_option("fade_linear", "Number of overlapping frames to crossfade: the end of each iteration is blended into the start of the next with a linear ramp."),
|
||||
_overlap_option("fade_smooth", "Number of overlapping frames to crossfade: the end of each iteration is blended into the start of the next with a smoothstep (ease in/out) ramp."),
|
||||
_overlap_option("start", "Number of frames to cut from the START of each iteration's output. Use when the model re-generates its context frames at the start of its output (most common for video continuation). The first iteration is trimmed only when initial_value is connected."),
|
||||
_overlap_option("end", "Number of frames to cut from the END of each iteration's output. Use when the model generates look-ahead frames at the end of its output that the next iteration re-generates. When initial_value is connected, the first iteration's start is also trimmed."),
|
||||
_overlap_option("fade_linear", "Number of overlapping frames to crossfade: the end of each iteration is blended into the start of the next with a linear ramp. When initial_value is connected, the first iteration's start is trimmed instead of blended."),
|
||||
_overlap_option("fade_smooth", "Number of overlapping frames to crossfade: the end of each iteration is blended into the start of the next with a smoothstep (ease in/out) ramp. When initial_value is connected, the first iteration's start is trimmed instead of blended."),
|
||||
]),
|
||||
io.Boolean.Input("stop", optional=True, default=False, raw_link=True, force_input=True,
|
||||
tooltip="Optional early stop signal from inside the loop body. When True, the loop stops after the current iteration regardless of remaining iterations or total_frames target."),
|
||||
@@ -164,7 +166,7 @@ class TensorLoopClose(io.ComfyNode):
|
||||
graph = GraphBuilder()
|
||||
open_id = flow_control[0]
|
||||
unpack = graph.node("_ImageAccumStateUnpack", loop_state=[open_id, 0])
|
||||
# unpack: 0=remaining, 1=accum, 2=previous_value, 3=accumulated_count, 4=count, 5=open_node_id, 6=total_frames
|
||||
# unpack: 0=remaining, 1=accum, 2=previous_value, 3=accumulated_count, 4=count, 5=open_node_id, 6=total_frames, 7=has_initial
|
||||
sub = graph.node("_IntOperations", operation="subtract", a=unpack.out(0), b=1)
|
||||
|
||||
overlap_mode = "disabled"
|
||||
@@ -177,26 +179,25 @@ class TensorLoopClose(io.ComfyNode):
|
||||
if accumulate:
|
||||
to_accum = processed
|
||||
if overlap_frames > 0 and overlap_mode != "disabled":
|
||||
# Iteration 1's start overlaps with initial_value (external context the model
|
||||
# re-generates), so it is only trimmed when initial_value was actually connected.
|
||||
# _BatchOps is a no-op at amount=0, so the trim amounts double as the conditionals.
|
||||
is_first = graph.node("_IntOperations", a=unpack.out(3), b=0, operation="==")
|
||||
trimmed_start = graph.node("_BatchOps", batch=processed, operation="trim_start", amount=overlap_frames).out(0)
|
||||
first_with_init = graph.node("_IntOperations", a=is_first.out(0), b=unpack.out(7), operation="multiply")
|
||||
|
||||
if overlap_mode == "start":
|
||||
to_accum = trimmed_start
|
||||
elif overlap_mode == "end":
|
||||
trimmed_end = graph.node("_BatchOps", batch=processed, operation="trim_end", amount=overlap_frames).out(0)
|
||||
trimmed_both = graph.node("_BatchOps", batch=trimmed_end, operation="trim_start", amount=overlap_frames).out(0)
|
||||
to_accum = graph.node("_ConditionalSelect",
|
||||
condition=is_first.out(1),
|
||||
value_if_true=trimmed_both,
|
||||
value_if_false=trimmed_end,
|
||||
).out(0)
|
||||
# Iterations 2+ always overlap the previous iteration's tail
|
||||
not_first = graph.node("_IntOperations", a=unpack.out(3), b=0, operation=">")
|
||||
do_trim = graph.node("_IntOperations", a=not_first.out(0), b=first_with_init.out(0), operation="add")
|
||||
else:
|
||||
# Fade: trim start on iter 1 only, keep full on subsequent for post-loop blend
|
||||
to_accum = graph.node("_ConditionalSelect",
|
||||
condition=is_first.out(1),
|
||||
value_if_true=trimmed_start,
|
||||
value_if_false=processed,
|
||||
).out(0)
|
||||
# end/fade keep iteration starts (end trims tails, fade blends post-loop)
|
||||
do_trim = first_with_init
|
||||
start_trim = graph.node("_IntOperations", a=do_trim.out(0), b=overlap_frames, operation="multiply")
|
||||
|
||||
to_trim = processed
|
||||
if overlap_mode == "end":
|
||||
to_trim = graph.node("_BatchOps", batch=processed, operation="trim_end", amount=overlap_frames).out(0)
|
||||
to_accum = graph.node("_BatchOps", batch=to_trim, operation="trim_start", amount=start_trim.out(0)).out(0)
|
||||
|
||||
accum_out = graph.node("_AccumulateNode", to_add=to_accum, accumulation=accum_out).out(0)
|
||||
|
||||
@@ -713,6 +714,7 @@ class _ImageAccumStateUnpack(io.ComfyNode):
|
||||
io.Int.Output("count"),
|
||||
io.AnyType.Output("open_node_id"),
|
||||
io.Int.Output("total_frames"),
|
||||
io.Int.Output("has_initial"),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -725,7 +727,8 @@ class _ImageAccumStateUnpack(io.ComfyNode):
|
||||
open_node_id = loop_state.get("open_node_id")
|
||||
total_frames = loop_state.get("total_frames", 0)
|
||||
accumulated_count = _accum_count(accum, loop_state.get("blend_overlap", 0))
|
||||
return io.NodeOutput(remaining, accum, previous_value, accumulated_count, count, open_node_id, total_frames)
|
||||
has_initial = int(bool(loop_state.get("has_initial", False)))
|
||||
return io.NodeOutput(remaining, accum, previous_value, accumulated_count, count, open_node_id, total_frames, has_initial)
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user