From 271e7de3001680aebe1cc0792710cab510563f55 Mon Sep 17 00:00:00 2001 From: Chris Dumas Date: Wed, 2 Sep 2026 12:55:51 +0000 Subject: [PATCH] Fix tag-driven H3 ref routing --- dumas_h3_longvideos.py | 53 +++++++++++++++++++++++-------- tests/test_dumas_h3_longvideos.py | 33 +++++++++++++++++++ 2 files changed, 72 insertions(+), 14 deletions(-) diff --git a/dumas_h3_longvideos.py b/dumas_h3_longvideos.py index a468a88..5bced8a 100644 --- a/dumas_h3_longvideos.py +++ b/dumas_h3_longvideos.py @@ -4635,6 +4635,31 @@ def resolve_prompt_refs(text, ref_list, include_named=True): return rewritten, refs, dropped +def resolve_tag_driven_prompt_refs(text, ref_list): + """Tag-driven prompt refs for mixed tagged/untagged scripts. + + Once any shot in the script uses ``, untagged shots should keep the + handoff path instead of quietly becoming ref-conditioned just because they name + a character or location. Named refs still supplement explicitly tagged shots so + their semantic context and actual image list stay aligned. + """ + rewritten, tagged_refs, dropped = resolve_tagged_refs(text, ref_list) + if not tagged_refs: + return rewritten, [], dropped + refs = list(tagged_refs) + seen_slots = { + slot_number + for slot_number in picture_tags(text) + if 1 <= slot_number <= len(ref_list or []) and _reference_image(ref_list[slot_number - 1]) is not None + } + for slot_number, ref in _named_refs_for_text(rewritten, ref_list, kinds=("character", "location")): + if slot_number in seen_slots: + continue + seen_slots.add(slot_number) + refs.append(ref) + return rewritten, refs, dropped + + def resolve_shot_references(text, ref_list, ref_mode="auto ref2v", shot_index=0, handoff=None): """Compatibility wrapper for older tests and helper code. @@ -6820,12 +6845,12 @@ class H3LongVideos: tagged_used = False on = [] effective_modes = [] - for shot_index, gen in enumerate(gens): - shot_mode = beat_ref_mode_directive(beats[shot_index] if shot_index < len(beats) else "") or ref_mode - if shot_mode in ("where tagged", "auto ref2v") and any_tags_anywhere: - if resolve_prompt_refs(gen, ref_slots)[1]: - on.append(shot_index + 1) - tagged_used = True + for shot_index, gen in enumerate(gens): + shot_mode = beat_ref_mode_directive(beats[shot_index] if shot_index < len(beats) else "") or ref_mode + if shot_mode in ("where tagged", "auto ref2v") and any_tags_anywhere: + if resolve_tag_driven_prompt_refs(gen, ref_slots)[1]: + on.append(shot_index + 1) + tagged_used = True else: mode_eff = ("every shot" if shot_mode == "auto ref2v" else "first shot" if shot_mode == "where tagged" else shot_mode) @@ -6914,14 +6939,14 @@ class H3LongVideos: shot_mode_eff = ("every shot" if shot_mode == "auto ref2v" else "first shot" if shot_mode == "where tagged" else shot_mode) carry_keyframe = False # tagged shot keeps its handoff as a keyframe - if shot_tag_driven: - # The prompt itself says where each reference belongs: the shot whose - # text names gets image N, renumbered to match what that - # shot actually carries. Every untagged shot keeps its handoff. - gen_prompt, shot_refs, dropped = resolve_prompt_refs(gen_prompt, ref_slots) - for n in dropped: - if n not in ref_missing: - ref_missing.append(n) + if shot_tag_driven: + # The prompt itself says where each reference belongs: the shot whose + # text names gets image N, renumbered to match what that + # shot actually carries. Every untagged shot keeps its handoff. + gen_prompt, shot_refs, dropped = resolve_tag_driven_prompt_refs(gen_prompt, ref_slots) + for n in dropped: + if n not in ref_missing: + ref_missing.append(n) else: shot_refs = shot_references(ref_slots, shot_mode_eff, i, handoff) # A shot that follows a strip starts FRESH. Continuing from a frame that diff --git a/tests/test_dumas_h3_longvideos.py b/tests/test_dumas_h3_longvideos.py index 36d52c3..3e80299 100644 --- a/tests/test_dumas_h3_longvideos.py +++ b/tests/test_dumas_h3_longvideos.py @@ -458,6 +458,39 @@ class DumasH3LongVideosHelperTests(unittest.TestCase): self.assertEqual([self.module._reference_image(ref) for ref in references], ["img2"]) self.assertEqual(dropped, []) + def test_resolve_tag_driven_prompt_refs_keeps_untagged_shot_on_handoff(self): + refs = [ + {"kind": "character", "image": "img1", "name": "Mara"}, + {"kind": "location", "image": "img2", "name": "Hangar"}, + ] + + text, references, dropped = self.module.resolve_tag_driven_prompt_refs( + "Mara waits in the hangar.", + refs, + ) + + self.assertEqual(text, "Mara waits in the hangar.") + self.assertEqual(references, []) + self.assertEqual(dropped, []) + + def test_resolve_tag_driven_prompt_refs_keeps_named_refs_on_tagged_shot(self): + refs = [ + {"kind": "character", "image": "img1", "name": "Mara"}, + {"kind": "location", "image": "img2", "name": "Hangar"}, + ] + + text, references, dropped = self.module.resolve_tag_driven_prompt_refs( + "Mara waits in the hangar near .", + refs, + ) + + self.assertEqual(text, "Mara waits in the hangar near .") + self.assertEqual( + [self.module._reference_image(ref) for ref in references], + ["img2", "img1"], + ) + self.assertEqual(dropped, []) + def test_resolve_shot_references_uses_named_characters_without_picture_tags(self): refs = [ {"kind": "character", "image": "img1", "name": "Mara"},