diff --git a/nodes.py b/nodes.py index 3dc88fe..32087a1 100644 --- a/nodes.py +++ b/nodes.py @@ -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)