Fix tag-driven H3 ref routing
This commit is contained in:
+39
-14
@@ -4635,6 +4635,31 @@ def resolve_prompt_refs(text, ref_list, include_named=True):
|
|||||||
return rewritten, refs, dropped
|
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 `<Picture N>`, 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):
|
def resolve_shot_references(text, ref_list, ref_mode="auto ref2v", shot_index=0, handoff=None):
|
||||||
"""Compatibility wrapper for older tests and helper code.
|
"""Compatibility wrapper for older tests and helper code.
|
||||||
|
|
||||||
@@ -6820,12 +6845,12 @@ class H3LongVideos:
|
|||||||
tagged_used = False
|
tagged_used = False
|
||||||
on = []
|
on = []
|
||||||
effective_modes = []
|
effective_modes = []
|
||||||
for shot_index, gen in enumerate(gens):
|
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
|
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 shot_mode in ("where tagged", "auto ref2v") and any_tags_anywhere:
|
||||||
if resolve_prompt_refs(gen, ref_slots)[1]:
|
if resolve_tag_driven_prompt_refs(gen, ref_slots)[1]:
|
||||||
on.append(shot_index + 1)
|
on.append(shot_index + 1)
|
||||||
tagged_used = True
|
tagged_used = True
|
||||||
else:
|
else:
|
||||||
mode_eff = ("every shot" if shot_mode == "auto ref2v"
|
mode_eff = ("every shot" if shot_mode == "auto ref2v"
|
||||||
else "first shot" if shot_mode == "where tagged" else shot_mode)
|
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"
|
shot_mode_eff = ("every shot" if shot_mode == "auto ref2v"
|
||||||
else "first shot" if shot_mode == "where tagged" else shot_mode)
|
else "first shot" if shot_mode == "where tagged" else shot_mode)
|
||||||
carry_keyframe = False # tagged shot keeps its handoff as a keyframe
|
carry_keyframe = False # tagged shot keeps its handoff as a keyframe
|
||||||
if shot_tag_driven:
|
if shot_tag_driven:
|
||||||
# The prompt itself says where each reference belongs: the shot whose
|
# The prompt itself says where each reference belongs: the shot whose
|
||||||
# text names <Picture N> gets image N, renumbered to match what that
|
# text names <Picture N> gets image N, renumbered to match what that
|
||||||
# shot actually carries. Every untagged shot keeps its handoff.
|
# shot actually carries. Every untagged shot keeps its handoff.
|
||||||
gen_prompt, shot_refs, dropped = resolve_prompt_refs(gen_prompt, ref_slots)
|
gen_prompt, shot_refs, dropped = resolve_tag_driven_prompt_refs(gen_prompt, ref_slots)
|
||||||
for n in dropped:
|
for n in dropped:
|
||||||
if n not in ref_missing:
|
if n not in ref_missing:
|
||||||
ref_missing.append(n)
|
ref_missing.append(n)
|
||||||
else:
|
else:
|
||||||
shot_refs = shot_references(ref_slots, shot_mode_eff, i, handoff)
|
shot_refs = shot_references(ref_slots, shot_mode_eff, i, handoff)
|
||||||
# A shot that follows a strip starts FRESH. Continuing from a frame that
|
# A shot that follows a strip starts FRESH. Continuing from a frame that
|
||||||
|
|||||||
@@ -458,6 +458,39 @@ class DumasH3LongVideosHelperTests(unittest.TestCase):
|
|||||||
self.assertEqual([self.module._reference_image(ref) for ref in references], ["img2"])
|
self.assertEqual([self.module._reference_image(ref) for ref in references], ["img2"])
|
||||||
self.assertEqual(dropped, [])
|
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 <Picture 2>.",
|
||||||
|
refs,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(text, "Mara waits in the hangar near <Picture 1>.")
|
||||||
|
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):
|
def test_resolve_shot_references_uses_named_characters_without_picture_tags(self):
|
||||||
refs = [
|
refs = [
|
||||||
{"kind": "character", "image": "img1", "name": "Mara"},
|
{"kind": "character", "image": "img1", "name": "Mara"},
|
||||||
|
|||||||
Reference in New Issue
Block a user