From 4c287491708af85f62e9e8479b7ebf1bbda369f3 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 15 Aug 2026 00:25:37 +0300 Subject: [PATCH] Initial qwen3/3.5/3.8 cuda graphs --- comfy/sd1_clip.py | 6 ++ comfy/text_encoders/gemma4.py | 3 - comfy/text_encoders/llama.py | 70 +++++++++++-- comfy/text_encoders/qwen35.py | 181 +++++++++++++++++----------------- 4 files changed, 163 insertions(+), 97 deletions(-) diff --git a/comfy/sd1_clip.py b/comfy/sd1_clip.py index f0fdf1aa5..8b25cf5a3 100644 --- a/comfy/sd1_clip.py +++ b/comfy/sd1_clip.py @@ -748,3 +748,9 @@ class SD1ClipModel(torch.nn.Module): def generate(self, tokens, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.95, min_p=0.0, repetition_penalty=1.0, seed=None, presence_penalty=0.0): return getattr(self, self.clip).generate(tokens, do_sample=do_sample, max_length=max_length, temperature=temperature, top_k=top_k, top_p=top_p, min_p=min_p, repetition_penalty=repetition_penalty, seed=seed, presence_penalty=presence_penalty) + + def get_dynamic_vram__units(self): + # forward to the inner transformer so ModelPatcher can register vbar units (graph decode) + model = getattr(getattr(getattr(self, self.clip), "transformer", None), "model", None) + get_units = getattr(model, "get_dynamic_vram__units", None) + return get_units() if get_units is not None else ([], []) diff --git a/comfy/text_encoders/gemma4.py b/comfy/text_encoders/gemma4.py index 61bc3a3f0..498f08df0 100644 --- a/comfy/text_encoders/gemma4.py +++ b/comfy/text_encoders/gemma4.py @@ -1719,9 +1719,6 @@ def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=No if dtype_llama is not None: dtype = dtype_llama super().__init__(device=device, dtype=dtype, name="gemma4", clip_model=clip_model, model_options=model_options) - - def get_dynamic_vram__units(self): - return getattr(self, self.clip).transformer.model.get_dynamic_vram__units() return Gemma4TEModel_ diff --git a/comfy/text_encoders/llama.py b/comfy/text_encoders/llama.py index f182e5147..1f14f2fe3 100644 --- a/comfy/text_encoders/llama.py +++ b/comfy/text_encoders/llama.py @@ -32,6 +32,31 @@ class FixedKV: def advance(self, num_tokens): self.index += num_tokens +@dataclass +class FixedKVBias(FixedKV): + # full-capacity decode bias [1, 1, 1, capacity]; position/bias/tracker are shared across layers + bias: torch.Tensor = None + tracker: dict = None + + def prepare(self, num_tokens): + if self.tracker["step"] == self.index: + return + if self.index + num_tokens > self.bias.shape[-1]: + raise RuntimeError("KV cache capacity exceeded") + self.tracker["step"] = self.index + self.position.fill_(self.index) + self.bias[..., self.index:self.index + num_tokens] = 0 + + +def fixed_kv_bias_decode(xq, cache, num_heads, num_kv_heads, head_dim): + # fixed-length masked attention over the full capacity, explicit math + batch_size = xq.shape[0] + groups = num_heads // num_kv_heads + q = xq.reshape(batch_size, num_kv_heads, groups, head_dim) * head_dim ** -0.5 + scores = (q @ cache.key.transpose(-1, -2)).add_(cache.bias) + probs = torch.softmax(scores, dim=-1, dtype=torch.float32).to(xq.dtype) + return (probs @ cache.value).reshape(batch_size, 1, num_heads * head_dim) + @dataclass class Llama2Config: vocab_size: int = 128320 @@ -278,6 +303,9 @@ class Qwen3VL_8BConfig(Qwen3_8BConfig): rope_theta: float = 5000000.0 rope_dims = [24, 20, 20] interleaved_mrope = True + fixed_kv: bool = True + graph_dynamic_vbar_blocks = True + prefetch_dynamic_vbars = True @dataclass class Qwen3VL_4BConfig(Qwen3VL_8BConfig): @@ -545,6 +573,7 @@ class Attention(nn.Module): optimized_attention=None, past_key_value: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, sliding_window: Optional[int] = None, + graph_decode: bool = False, ): batch_size, seq_length, _ = hidden_states.shape @@ -567,7 +596,18 @@ class Attention(nn.Module): xq, xk = apply_rope(xq, xk, freqs_cis=freqs_cis) fixed_cache = past_key_value if isinstance(past_key_value, FixedKV) else None - if fixed_cache is not None: + if isinstance(fixed_cache, FixedKVBias): + if seq_length == 1 and attention_mask is None and graph_decode: + # CUDA-graphable decode: device-side write position, full-capacity biased attention + fixed_cache.key.index_copy_(2, fixed_cache.position, xk) + fixed_cache.value.index_copy_(2, fixed_cache.position, xv) + output = fixed_kv_bias_decode(xq, fixed_cache, self.num_heads, self.num_kv_heads, self.head_dim) + return self.o_proj(output), fixed_cache + fixed_cache.key[:, :, fixed_cache.index:fixed_cache.index + seq_length] = xk + fixed_cache.value[:, :, fixed_cache.index:fixed_cache.index + seq_length] = xv + xk = fixed_cache.key[:, :, :fixed_cache.index + seq_length] + xv = fixed_cache.value[:, :, :fixed_cache.index + seq_length] + elif fixed_cache is not None: xq = xq.transpose(1, 2) xk = xk.transpose(1, 2) xv = xv.transpose(1, 2) @@ -659,6 +699,12 @@ class TransformerBlock(nn.Module): past_key_value: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, ): output = x + # mirror the conditions under which prefetch_queue_pop can actually capture, so + # eager fallbacks keep the sliced decode path instead of the full-capacity one + graph_decode = (getattr(self, "_v_block", None) is not None + and comfy.model_management.NUM_STREAMS > 0 + and not comfy.model_management.args.disable_cuda_graphs + and comfy.model_management.is_device_cuda(x.device)) # Self Attention residual = x x = self.input_layernorm(x) @@ -668,6 +714,7 @@ class TransformerBlock(nn.Module): freqs_cis=freqs_cis, optimized_attention=optimized_attention, past_key_value=past_key_value, + graph_decode=graph_decode, ) x = residual + x @@ -759,7 +806,8 @@ class Llama2_(nn.Module): super().__init__() self.config = config self.fixed_kv = getattr(config, "fixed_kv", False) - self.graph_dynamic_vbar_blocks = False + self.graph_dynamic_vbar_blocks = getattr(config, "graph_dynamic_vbar_blocks", False) + self.prefetch_dynamic_vbars = getattr(config, "prefetch_dynamic_vbars", False) self.vocab_size = config.vocab_size if self.config.transformer_type == "gemma2" or self.config.transformer_type == "gemma3": @@ -791,14 +839,24 @@ class Llama2_(nn.Module): def init_kv_cache(self, batch, capacity, device, dtype): caches = [] - fixed_kv = self.fixed_kv and comfy_kitchen.flash_attention_decode_is_available(device) + flash = getattr(comfy_kitchen, "flash_attention_decode_is_available", None) + flash_kv = self.fixed_kv and flash is not None and flash(device) + bias_kv = self.fixed_kv and not flash_kv + # all layers advance in lockstep, so the bias caches share one position/bias/tracker + position = torch.empty((1,), device=device, dtype=torch.int64) if bias_kv else None + bias = torch.full((1, 1, 1, capacity), torch.finfo(dtype).min, device=device, dtype=dtype) if bias_kv else None + tracker = {"step": -1} for _ in range(self.config.num_hidden_layers): - if fixed_kv: + if flash_kv: key = torch.empty((batch, capacity, self.config.num_key_value_heads, self.config.head_dim), device=device, dtype=dtype) value = torch.empty_like(key) - position = torch.empty((1,), device=device, dtype=torch.int64) + pos = torch.empty((1,), device=device, dtype=torch.int64) seqlen = torch.empty((batch,), device=device, dtype=torch.int32) - caches.append(FixedKV(key, value, 0, position, seqlen)) + caches.append(FixedKV(key, value, 0, pos, seqlen)) + elif bias_kv: + # zero-init: decode attends full capacity with masked tails, 0*0 stays finite + key = torch.zeros((batch, self.config.num_key_value_heads, capacity, self.config.head_dim), device=device, dtype=dtype) + caches.append(FixedKVBias(key, torch.zeros_like(key), 0, position, None, bias, tracker)) else: key = torch.empty((batch, self.config.num_key_value_heads, capacity, self.config.head_dim), device=device, dtype=dtype) caches.append((key, torch.empty_like(key), 0)) diff --git a/comfy/text_encoders/qwen35.py b/comfy/text_encoders/qwen35.py index 304a4357f..544c64e21 100644 --- a/comfy/text_encoders/qwen35.py +++ b/comfy/text_encoders/qwen35.py @@ -5,11 +5,32 @@ from dataclasses import dataclass, field import os import comfy.model_management +import comfy.ops from comfy.ldm.modules.attention import optimized_attention_for_device from comfy import sd1_clip import comfy.text_encoders.qwen_vl -from .llama import BaseLlama, BaseGenerate, Llama2_, MLP, RMSNorm, apply_rope +from .llama import BaseLlama, BaseGenerate, FixedKV, FixedKVBias, Llama2_, MLP, RMSNorm, apply_rope, fixed_kv_bias_decode, precompute_freqs_cis + + +@dataclass +class LinearKV(FixedKV): + # DeltaNet state on the FixedKV interface: key=conv_state, value=recurrent_state (fp32) + g_decay: torch.Tensor = None + dt_bias: torch.Tensor = None + + def prepare(self, num_tokens): + pass + + @property + def conv_state(self): + return self.key + + @property + def recurrent_state(self): + return self.value + + def _qwen35_layer_types(n): @@ -135,18 +156,6 @@ def torch_chunk_gated_delta_rule(query, key, value, g, beta, chunk_size=64, init return core_attn_out, last_recurrent_state -def torch_causal_conv1d_update(x, conv_state, weight, bias=None): - # conv_state: [B, channels, kernel_size-1], x: [B, channels, 1] - # weight: [channels, kernel_size] - state_len = conv_state.shape[-1] - combined = torch.cat([conv_state, x], dim=-1).to(weight.dtype) # [B, channels, kernel_size] - conv_state.copy_(combined[:, :, -state_len:]) - out = (combined * weight).sum(dim=-1, keepdim=True) # [B, channels, 1] - if bias is not None: - out = out + bias.unsqueeze(0).unsqueeze(-1) - return F.silu(out).to(x.dtype) - - # GatedDeltaNet - Linear Attention Layer class GatedDeltaNet(nn.Module): @@ -185,7 +194,7 @@ class GatedDeltaNet(nn.Module): use_recurrent = ( past_key_value is not None - and past_key_value[2] > 0 + and past_key_value.index > 0 and seq_len == 1 ) @@ -197,13 +206,14 @@ class GatedDeltaNet(nn.Module): # Conv1d if use_recurrent: - recurrent_state, conv_state, step_index = past_key_value - conv_weight = comfy.model_management.cast_to_device(self.conv1d.weight, mixed_qkv.device, mixed_qkv.dtype).squeeze(1) - conv_bias = comfy.model_management.cast_to_device(self.conv1d.bias, mixed_qkv.device, mixed_qkv.dtype) if self.conv1d.bias is not None else None - mixed_qkv = torch_causal_conv1d_update(mixed_qkv, conv_state, conv_weight, conv_bias) + # decode: exact-width causal window, weight resolved via the vbar-aware context + combined = torch.cat([past_key_value.conv_state, mixed_qkv], dim=-1) + past_key_value.conv_state.copy_(combined[:, :, 1:]) + with comfy.ops.CastBiasWeightContext(self.conv1d, combined, offloadable=True) as (conv_weight, conv_bias): + mixed_qkv = F.silu(F.conv1d(combined, conv_weight, conv_bias, groups=self.conv1d.groups)) else: if past_key_value is not None: - recurrent_state, conv_state, step_index = past_key_value + conv_state = past_key_value.conv_state conv_state_init = F.pad(mixed_qkv, (self.conv_kernel_size - mixed_qkv.shape[-1], 0)) conv_state.copy_(conv_state_init[:, :, -conv_state.shape[-1]:]) mixed_qkv = F.silu(self.conv1d(mixed_qkv)[:, :, :seq_len]) @@ -212,7 +222,15 @@ class GatedDeltaNet(nn.Module): mixed_qkv = mixed_qkv.transpose(1, 2) # [B, seq_len, conv_dim] query, key, value = mixed_qkv.split([self.key_dim, self.key_dim, self.value_dim], dim=-1) beta = b.sigmoid() - g = -self.A_log.float().exp() * F.softplus(a.float() + self.dt_bias.float()) + if use_recurrent: + g_decay, dt_bias = past_key_value.g_decay, past_key_value.dt_bias + else: + g_decay = -comfy.model_management.cast_to_device(self.A_log, x.device, torch.float32).exp() + dt_bias = comfy.model_management.cast_to_device(self.dt_bias, x.device, torch.float32) + if past_key_value is not None: + past_key_value.g_decay = g_decay + past_key_value.dt_bias = dt_bias + g = g_decay * F.softplus(a.float() + dt_bias) # Delta rule if use_recurrent: @@ -234,6 +252,7 @@ class GatedDeltaNet(nn.Module): g_t = g.reshape(batch_size, -1).exp() # In-place state update: [B, heads, k_dim, v_dim] + recurrent_state = past_key_value.recurrent_state recurrent_state.mul_(g_t[:, :, None, None]) kv_mem = torch.einsum('bhk,bhkv->bhv', k, recurrent_state) delta = (v - kv_mem) * beta_t[:, :, None] @@ -241,7 +260,7 @@ class GatedDeltaNet(nn.Module): core_attn_out = torch.einsum('bhk,bhkv->bhv', q, recurrent_state) core_attn_out = core_attn_out.to(x.dtype).unsqueeze(1) - present_key_value = (recurrent_state, conv_state, step_index + 1) + present_key_value = past_key_value else: query = query.reshape(batch_size, seq_len, -1, self.key_head_dim) key = key.reshape(batch_size, seq_len, -1, self.key_head_dim) @@ -261,8 +280,8 @@ class GatedDeltaNet(nn.Module): present_key_value = None if past_key_value is not None: if last_recurrent_state is not None: - recurrent_state.copy_(last_recurrent_state.to(recurrent_state.dtype)) - present_key_value = (recurrent_state, conv_state, step_index + seq_len) + past_key_value.recurrent_state.copy_(last_recurrent_state.to(past_key_value.recurrent_state.dtype)) + present_key_value = past_key_value # Gated norm + output projection (shared) core_attn_out = self.norm(core_attn_out.reshape(-1, self.value_head_dim), z.reshape(-1, self.value_head_dim)) @@ -271,29 +290,6 @@ class GatedDeltaNet(nn.Module): # GatedAttention - Full Attention with output gating -def precompute_partial_rope(head_dim, rotary_dim, position_ids, theta, device=None, mrope_section=None): - """Compute RoPE frequencies for partial rotary embeddings.""" - theta_numerator = torch.arange(0, rotary_dim, 2, device=device).float() - inv_freq = 1.0 / (theta ** (theta_numerator / rotary_dim)) - - inv_freq_expanded = inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) - position_ids_expanded = position_ids[:, None, :].float() - freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) - emb = torch.cat((freqs, freqs), dim=-1) - cos = emb.cos() - sin = emb.sin() - - if mrope_section is not None and position_ids.shape[0] == 3: - mrope_section_2 = [s * 2 for s in mrope_section] - cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section_2, dim=-1))], dim=-1).unsqueeze(0) - sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section_2, dim=-1))], dim=-1).unsqueeze(0) - - cos = cos.unsqueeze(1) - sin = sin.unsqueeze(1) - sin_split = sin.shape[-1] // 2 - return (cos, sin[..., :sin_split], -sin[..., sin_split:]) - - def apply_partial_rope(xq, xk, freqs_cis, rotary_dim): """Apply RoPE to only the first rotary_dim dimensions.""" xq_rot = xq[..., :rotary_dim] @@ -329,7 +325,7 @@ class GatedAttention(nn.Module): self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps, add=config.rms_norm_add, device=device, dtype=dtype) self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps, add=config.rms_norm_add, device=device, dtype=dtype) - def forward(self, x, attention_mask=None, freqs_cis=None, optimized_attention=None, past_key_value=None): + def forward(self, x, attention_mask=None, freqs_cis=None, optimized_attention=None, past_key_value=None, graph_decode=False): batch_size, seq_length, _ = x.shape # Project Q (with gate), K, V @@ -350,26 +346,24 @@ class GatedAttention(nn.Module): xq, xk = apply_partial_rope(xq, xk, freqs_cis, self.rotary_dim) # KV cache - present_key_value = None - if past_key_value is not None: - past_key, past_value, index = past_key_value - num_tokens = xk.shape[2] - if past_key.shape[2] >= (index + num_tokens): - past_key[:, :, index:index + num_tokens] = xk - past_value[:, :, index:index + num_tokens] = xv - xk = past_key[:, :, :index + num_tokens] - xv = past_value[:, :, :index + num_tokens] - present_key_value = (past_key, past_value, index + num_tokens) - else: - if index > 0: - xk = torch.cat((past_key[:, :, :index], xk), dim=2) - xv = torch.cat((past_value[:, :, :index], xv), dim=2) - present_key_value = (xk, xv, index + num_tokens) + present_key_value = past_key_value + if past_key_value is not None and seq_length == 1 and attention_mask is None and graph_decode: + # CUDA-graphable decode: device-side write position, full-capacity biased attention + cache = past_key_value + cache.key.index_copy_(2, cache.position, xk) + cache.value.index_copy_(2, cache.position, xv) + output = fixed_kv_bias_decode(xq, cache, self.num_heads, self.num_kv_heads, self.head_dim) + else: + if past_key_value is not None: + cache = past_key_value + cache.key[:, :, cache.index:cache.index + seq_length] = xk + cache.value[:, :, cache.index:cache.index + seq_length] = xv + xk = cache.key[:, :, :cache.index + seq_length] + xv = cache.value[:, :, :cache.index + seq_length] + gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} + output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True, **gqa_kwargs) - gqa_kwargs = {"enable_gqa": True} if self.num_heads != self.num_kv_heads else {} - output = optimized_attention(xq, xk, xv, self.num_heads, mask=attention_mask, skip_reshape=True, **gqa_kwargs) output = output * gate.sigmoid() - return self.o_proj(output), present_key_value @@ -387,13 +381,21 @@ class Qwen35TransformerBlock(nn.Module): self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, add=config.rms_norm_add, device=device, dtype=dtype) def forward(self, x, attention_mask=None, freqs_cis=None, optimized_attention=None, past_key_value=None): + output = x if self.layer_type == "linear_attention": h, present_key_value = self.linear_attn(self.input_layernorm(x), attention_mask=attention_mask, past_key_value=past_key_value) else: - h, present_key_value = self.self_attn(self.input_layernorm(x), attention_mask=attention_mask, freqs_cis=freqs_cis, optimized_attention=optimized_attention, past_key_value=past_key_value) + # mirror the conditions under which prefetch_queue_pop can actually capture, so + # eager fallbacks keep the sliced decode path instead of the full-capacity one + graph_decode = (getattr(self, "_v_block", None) is not None + and comfy.model_management.NUM_STREAMS > 0 + and not comfy.model_management.args.disable_cuda_graphs + and comfy.model_management.is_device_cuda(x.device)) + h, present_key_value = self.self_attn(self.input_layernorm(x), attention_mask=attention_mask, freqs_cis=freqs_cis, optimized_attention=optimized_attention, past_key_value=past_key_value, graph_decode=graph_decode) - x = x + h - x = x + self.mlp(self.post_attention_layernorm(x)) + # in-place into the input buffer so CUDA-graph replays land in the static x + x = torch.add(x, h, out=output) + x = torch.add(x, self.mlp(self.post_attention_layernorm(x)), out=output) return x, present_key_value @@ -402,6 +404,8 @@ class Qwen35Transformer(Llama2_): def __init__(self, config, device=None, dtype=None, ops=None): nn.Module.__init__(self) self.config = config + self.prefetch_dynamic_vbars = True + self.graph_dynamic_vbar_blocks = True self.vocab_size = config.vocab_size self.embed_tokens = ops.Embedding(config.vocab_size, config.hidden_size, device=device, dtype=dtype) self.layers = nn.ModuleList([ @@ -417,21 +421,10 @@ class Qwen35Transformer(Llama2_): if config.lm_head: self.lm_head = ops.Linear(config.hidden_size, config.vocab_size, bias=False, device=device, dtype=dtype) - def get_past_len(self, past_key_values): - for i, layer in enumerate(self.layers): - if layer.layer_type == "full_attention": - if len(past_key_values) > i: - return past_key_values[i][2] - break - return 0 - def compute_freqs_cis(self, position_ids, device): rotary_dim = int(self.config.head_dim * self.config.partial_rotary_factor) - return precompute_partial_rope( - self.config.head_dim, rotary_dim, position_ids, - self.config.rope_theta, device=device, - mrope_section=self.config.mrope_section, - ) + return precompute_freqs_cis(rotary_dim, position_ids, self.config.rope_theta, + rope_dims=self.config.mrope_section, interleaved_mrope=True, device=device) # Vision Encoder @@ -687,7 +680,8 @@ class Qwen35(BaseLlama, BaseGenerate, torch.nn.Module): def preprocess_embed(self, embed, device): if embed["type"] == "image": - image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images(embed["data"], patch_size=16) + # Qwen3.5 normalizes to [-1, 1] (mean/std 0.5), same as Qwen3-VL. + image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images(embed["data"], patch_size=16, image_mean=[0.5, 0.5, 0.5], image_std=[0.5, 0.5, 0.5]) return self.visual(image.to(device, dtype=torch.float32), grid), grid return None, None @@ -698,6 +692,10 @@ class Qwen35(BaseLlama, BaseGenerate, torch.nn.Module): def init_kv_cache(self, batch, max_cache_len, device, execution_dtype): model_config = self.model.config past_key_values = [] + # all full-attention layers advance in lockstep, so they share one position/bias/tracker + position = torch.empty((1,), device=device, dtype=torch.int64) + bias = torch.full((1, 1, 1, max_cache_len), torch.finfo(execution_dtype).min, device=device, dtype=execution_dtype) + tracker = {"step": -1} for i in range(model_config.num_hidden_layers): if model_config.layer_types[i] == "linear_attention": recurrent_state = torch.zeros( @@ -709,13 +707,11 @@ class Qwen35(BaseLlama, BaseGenerate, torch.nn.Module): [batch, conv_dim, model_config.conv_kernel_size - 1], device=device, dtype=execution_dtype ) - past_key_values.append((recurrent_state, conv_state, 0)) + past_key_values.append(LinearKV(conv_state, recurrent_state, 0, None, None)) else: - past_key_values.append(( - torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), - torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), - 0 - )) + # zero-init: decode attends full capacity with masked tails, 0*0 stays finite + key = torch.zeros([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype) + past_key_values.append(FixedKVBias(key, torch.zeros_like(key), 0, position, None, bias, tracker)) return past_key_values # Tokenizer and Text Encoder Wrappers @@ -786,6 +782,15 @@ class Qwen35ClipModel(sd1_clip.SDClipModel): dtype=dtype, special_tokens={"pad": 248044}, layer_norm_hidden_state=False, model_class=Qwen35_, enable_attention_masks=attention_mask, return_attention_masks=attention_mask, model_options=model_options) + def generate(self, tokens, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, presence_penalty=0.0): + if isinstance(tokens, dict): + tokens = next(iter(tokens.values())) + tokens_only = [[t[0] for t in b] for b in tokens] + embeds, _, _, embeds_info = self.process_tokens(tokens_only, self.execution_device) + position_ids = comfy.text_encoders.qwen_vl.qwen2vl_mrope_position_ids(embeds_info, embeds.shape[1], embeds.device) + return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, + presence_penalty=presence_penalty, position_ids=position_ids) + class Qwen35TEModel(sd1_clip.SD1ClipModel): def __init__(self, device="cpu", dtype=None, model_options={}, model_type="qwen35_2b"):