From c2aeaa2c177e0c760e332ecf73b52dc60ba730b3 Mon Sep 17 00:00:00 2001 From: silveroxides Date: Sun, 19 Jul 2026 10:37:57 +0200 Subject: [PATCH] fix: use Krea template for image fusion Restore the original Krea prefix stripping logic and place visual tokens after the user prefix. --- comfy/text_encoders/krea2.py | 30 ++++----- .../test_qwen_visual_fusion.py | 66 ++++++++++++++++--- 2 files changed, 71 insertions(+), 25 deletions(-) diff --git a/comfy/text_encoders/krea2.py b/comfy/text_encoders/krea2.py index 0341dd6e3..db8427c7c 100644 --- a/comfy/text_encoders/krea2.py +++ b/comfy/text_encoders/krea2.py @@ -20,20 +20,6 @@ KREA2_TAP_LAYERS = [2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35] KREA2_TEMPLATE = "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" -def _krea2_template_end(tok_pairs): - image_position = next((i for i, pair in enumerate(tok_pairs) if isinstance(pair[0], dict) and pair[0].get("type") == "image"), len(tok_pairs)) - template_end = -1 - for i in range(image_position): - if i + 2 >= len(tok_pairs): - break - values = [tok_pairs[j][0] for j in range(i, i + 3)] - if all(not torch.is_tensor(value) and isinstance(value, numbers.Integral) for value in values) and values == [151644, 872, 198]: - template_end = i + 3 - if template_end == -1: - raise ValueError("Could not locate the Krea 2 user prompt template.") - return template_end - - class Krea2Tokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer): def __init__(self, embedding_directory=None, tokenizer_data={}): super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, model_type="qwen3vl_4b") @@ -41,6 +27,9 @@ class Krea2Tokenizer(comfy.text_encoders.qwen3vl.Qwen3VLTokenizer): def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=[], prevent_empty_text=False, thinking=True, **kwargs): # Krea2 conditions on the no-think template; thinking=True drops the empty block qwen3vl adds. + if llama_template is None and len(images) > 0 and not text.startswith('<|im_start|>'): + vision_block = "<|vision_start|><|image_pad|><|vision_end|>" + llama_template = KREA2_TEMPLATE.replace("{}", vision_block * len(images) + "{}") return super().tokenize_with_weights(text, return_word_ids=return_word_ids, llama_template=llama_template, images=images, prevent_empty_text=prevent_empty_text, thinking=thinking, **kwargs) @@ -58,8 +47,19 @@ class Krea2TEModel(sd1_clip.SD1ClipModel): out, pooled, extra = super().encode_token_weights(token_weight_pairs) # out: (B, 12, seq, 2560) tok_pairs = token_weight_pairs["qwen3vl_4b"][0] + # Strip the system + user-opening prefix + count_im_start = 0 if template_end == -1: - template_end = _krea2_template_end(tok_pairs) + for i, v in enumerate(tok_pairs): + elem = v[0] + if not torch.is_tensor(elem) and isinstance(elem, numbers.Integral): + if elem == 151644 and count_im_start < 2: + template_end = i + count_im_start += 1 + if out.shape[2] > (template_end + 3): + if tok_pairs[template_end + 1][0] == 872: # "user" + if tok_pairs[template_end + 2][0] == 198: # "\n" + template_end += 3 out = out[:, :, template_end:] diff --git a/tests-unit/comfy_extras_test/test_qwen_visual_fusion.py b/tests-unit/comfy_extras_test/test_qwen_visual_fusion.py index ff3c104c7..1b37925ca 100644 --- a/tests-unit/comfy_extras_test/test_qwen_visual_fusion.py +++ b/tests-unit/comfy_extras_test/test_qwen_visual_fusion.py @@ -8,7 +8,7 @@ if not torch.cuda.is_available(): cli_args.cpu = True try: - from comfy.text_encoders.krea2 import KREA2_TEMPLATE, Krea2Tokenizer, _krea2_template_end + from comfy.text_encoders.krea2 import KREA2_TEMPLATE, Krea2Tokenizer from comfy_extras.nodes_cond import CLIPTextEncodeImageFusion, _flatten_images, _fuse_conditionings, _resize_visual_tokens, _spatial_fusion_mask, _visual_grid, _visual_token_span finally: cli_args.cpu = prior_cpu @@ -63,20 +63,66 @@ def test_visual_span_accounts_for_stripped_prefix(): assert _visual_token_span(tokens, cond_length=9, visual_tokens=4) == (1, 5) -def test_krea_template_stripping_preserves_visual_tokens(): +def test_krea_conditioning_template_preserves_visual_tokens(): + tokenizer = Krea2Tokenizer() + image = torch.zeros(1, 64, 64, 3) + + for image_count in (1, 2, 3): + tokens = tokenizer.tokenize_with_weights("test prompt", images=[image] * image_count) + token_pairs = tokens["qwen3vl_4b"][0] + image_positions = [i for i, pair in enumerate(token_pairs) if isinstance(pair[0], dict)] + im_start_positions = [i for i, pair in enumerate(token_pairs) if pair[0] == 151644] + template_end = im_start_positions[1] + 3 + + assert len(image_positions) == image_count + assert token_pairs[im_start_positions[1] + 1][0] == 872 + assert token_pairs[im_start_positions[1] + 2][0] == 198 + assert template_end <= image_positions[0] + + if image_count == 1: + cond_length = len(token_pairs) - 1 + 4 - template_end + assert _visual_token_span(tokens, cond_length, 4) == (image_positions[0] - template_end, image_positions[0] - template_end + 4) + + +def test_krea_preformatted_and_text_generation_templates_are_unchanged(): tokenizer = Krea2Tokenizer() image = torch.zeros(1, 64, 64, 3) image_prompt = "<|vision_start|><|image_pad|><|vision_end|>test prompt" - for text in ("test prompt", KREA2_TEMPLATE.format(image_prompt)): - tokens = tokenizer.tokenize_with_weights(text, images=[image]) - token_pairs = tokens["qwen3vl_4b"][0] - image_position = next(i for i, pair in enumerate(token_pairs) if isinstance(pair[0], dict)) - template_end = _krea2_template_end(token_pairs) - cond_length = len(token_pairs) - 1 + 4 - template_end + preformatted = tokenizer.tokenize_with_weights(KREA2_TEMPLATE.format(image_prompt), images=[image])["qwen3vl_4b"][0] + generated = tokenizer.tokenize_with_weights("test prompt", image=image, thinking=False)["qwen3vl_4b"][0] - assert template_end <= image_position - assert _visual_token_span(tokens, cond_length, 4) == (image_position - template_end, image_position - template_end + 4) + assert [pair[0] for pair in preformatted[:3]] == [151644, 8948, 198] + assert [pair[0] for pair in generated[:3]] == [151644, 872, 198] + + +def test_krea_tokenization_runs_through_image_fusion_node(): + class FakeKreaClip: + def __init__(self): + self.tokenizer = Krea2Tokenizer() + + def tokenize(self, text, images): + return self.tokenizer.tokenize_with_weights(text, images=images) + + def encode_from_tokens_scheduled(self, tokens): + token_pairs = tokens["qwen3vl_4b"][0] + im_start_positions = [i for i, pair in enumerate(token_pairs) if pair[0] == 151644] + values = [] + for pair in token_pairs[im_start_positions[1] + 3:]: + if isinstance(pair[0], dict): + values.extend([pair[0]["data"].mean()] * 4) + else: + values.append(torch.tensor(0.0)) + return [[torch.stack(values).reshape(1, -1, 1), {}]] + + images = { + "image_1": torch.zeros(1, 64, 64, 3), + "image_2": torch.ones(1, 64, 64, 3), + } + + conditioning = CLIPTextEncodeImageFusion.execute(FakeKreaClip(), "test prompt", images, "spatial-checkerboard").args[0] + + assert set(conditioning[0][0][:, 1:5].flatten().tolist()) == {0.0, 1.0} def test_fusion_replaces_only_visual_tokens_and_preserves_dtype_and_metadata():