diff --git a/dumas_h3_latent_upscale.py b/dumas_h3_latent_upscale.py index 624af60..f2941ba 100644 --- a/dumas_h3_latent_upscale.py +++ b/dumas_h3_latent_upscale.py @@ -16,6 +16,11 @@ except Exception: # pragma: no cover - import-time fallback for the test shim nn = None F = None +try: + from comfy.ldm.minimax.model import FRAME_PER_TOKEN +except Exception: # pragma: no cover - import-time fallback for the test shim + FRAME_PER_TOKEN = (1, 4, 4, 4, 4) + LATENTS_MEAN = [ 0.858090341091156, -0.9606591463088989, 1.0661640167236328, -0.5090325474739075, @@ -372,30 +377,63 @@ def _resolve_target_size(param, h_in, w_in): return int(h_in), int(w_in) -def _temporal_segments(frame_count, chunk_length, overlap): - frame_count = int(frame_count) +def _frames_for_tokens(n): + return sum(FRAME_PER_TOKEN[i % 5] for i in range(int(n))) + + +def _snap_frame_boundary(f, max_tokens, phase=5): + best_k, best_f, best_d = 0, 0, abs(int(f)) + for k in range(0, int(max_tokens) + 1, phase): + acc = _frames_for_tokens(k) + d = abs(acc - f) + if d < best_d: + best_k, best_f, best_d = k, acc, d + return best_k, best_f + + +def _temporal_segments(token_count, chunk_length, overlap): + token_count = int(token_count) chunk_length = int(chunk_length) overlap = int(overlap) - if frame_count <= 0: - return [(0, 0)] + if token_count <= 0: + return [(0, 0, 0, 0)] if chunk_length <= 0: raise ValueError("chunk_length must be positive") if overlap < 0: raise ValueError("temporal_overlap must be non-negative") if chunk_length <= overlap: raise ValueError("temporal_overlap must be smaller than chunk_length") + frame_count = _frames_for_tokens(token_count) if frame_count <= chunk_length: - return [(0, frame_count)] + return [(0, 0, token_count, frame_count)] hop = chunk_length - overlap bounds = [] - start = 0 - while start < frame_count: - end = min(start + chunk_length, frame_count) - bounds.append((start, end)) - if end >= frame_count: + prev_end_k = 0 + i = 0 + while True: + s = i * hop + e = min(s + chunk_length, frame_count) + if i == 0: + k0, f0 = 0, 0 + else: + k0, f0 = _snap_frame_boundary(s, token_count, phase=5) + if k0 > prev_end_k: + k0, f0 = prev_end_k, _frames_for_tokens(prev_end_k) + if e >= frame_count: + k1, f1 = token_count, frame_count + else: + k1, f1 = _snap_frame_boundary(e, token_count, phase=5) + if k1 <= k0: + k1 = min(token_count, k0 + 5) + f1 = _frames_for_tokens(k1) + if k1 >= token_count: + k1, f1 = token_count, frame_count + bounds.append((k0, f0, k1, f1)) + if k1 >= token_count: break - start += hop + prev_end_k = k1 + i += 1 return bounds @@ -712,7 +750,8 @@ def _upscale_video_temporal_chunks(video, param, upscaler): temporal_overlap = int(param.get("temporal_overlap", 0) or 0) anchor_strength = float(param.get("anchor_strength", 0.999) or 0.999) t = int(video.shape[2]) - if chunk_length <= 0 or t <= chunk_length: + frame_count = _frames_for_tokens(t) + if chunk_length <= 0 or frame_count <= chunk_length: return upscaler(video, param) bounds = _temporal_segments(t, chunk_length, temporal_overlap) @@ -722,8 +761,8 @@ def _upscale_video_temporal_chunks(video, param, upscaler): orig_dtype = video.dtype out = None out_h = out_w = None - for i, (t0, t1) in enumerate(bounds): - chunk = video[:, :, t0:t1].contiguous() + for i, (k0, f0, k1, f1) in enumerate(bounds): + chunk = video[:, :, k0:k1].contiguous() try: chunk_out, chunk_h, chunk_w = upscaler(chunk, param) except RuntimeError as exc: @@ -746,22 +785,22 @@ def _upscale_video_temporal_chunks(video, param, upscaler): if out is None: out_h, out_w = chunk_h, chunk_w out = torch.zeros((video.shape[0], video.shape[1], t, out_h, out_w), device="cpu", dtype=orig_dtype) - out[:, :, t0:t1] = chunk_out + out[:, :, k0:k1] = chunk_out continue - ov = min(temporal_overlap, t1 - t0, t - t0) + ov = min(temporal_overlap, k1 - k0, t - k0) if ov <= 0: - out[:, :, t0:t1] = chunk_out + out[:, :, k0:k1] = chunk_out continue - prev_region = out[:, :, t0:t0 + ov] + prev_region = out[:, :, k0:k0 + ov] new_region = chunk_out[:, :, :ov] w = _temporal_blend_weights(ov, anchor_strength).to(device=prev_region.device, dtype=prev_region.dtype) - out[:, :, t0:t0 + ov] = ( + out[:, :, k0:k0 + ov] = ( prev_region * (1.0 - w[None, None, :, None, None]) + new_region * w[None, None, :, None, None] ) - out[:, :, t0 + ov:t1] = chunk_out[:, :, ov:] + out[:, :, k0 + ov:k1] = chunk_out[:, :, ov:] return out, out_h, out_w diff --git a/tests/test_dumas_h3_longvideos.py b/tests/test_dumas_h3_longvideos.py index d7956d1..3d456e3 100644 --- a/tests/test_dumas_h3_longvideos.py +++ b/tests/test_dumas_h3_longvideos.py @@ -1181,8 +1181,10 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): def test_temporal_segments_split_long_sequences(self): latent = importlib.import_module("dumas_h3_latent_upscale") - bounds = latent._temporal_segments(124, 85, 17) - self.assertEqual(bounds, [(0, 85), (68, 124)]) + bounds = latent._temporal_segments(36, 85, 17) + self.assertGreater(len(bounds), 1) + self.assertEqual(bounds[0][0], 0) + self.assertEqual(bounds[-1][2], 36) def test_shrink_temporal_param_reduces_chunk_length(self): latent = importlib.import_module("dumas_h3_latent_upscale")