Skip to content

Commit 1b01a6c

Browse files
Merge branch 'master' into jaeone94/native-lora-stack
2 parents 59e1c3c + a4b5a04 commit 1b01a6c

26 files changed

Lines changed: 1016 additions & 157 deletions

File tree

‎.coderabbit.yaml‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -140,5 +140,16 @@ knowledge_base:
140140
filePatterns:
141141
- files: "AGENTS.md"
142142
applyTo: "**"
143+
linked_repositories:
144+
- repository: "Comfy-Org/ComfyUI_frontend"
145+
instructions: |
146+
Frontend consumer of Core APIs, routes, feature flags, asset behavior,
147+
node metadata, and package-version rollouts. Check for compatibility
148+
issues and changes that require coordinated rollout or merge ordering.
149+
- repository: "Comfy-Org/comfy-kitchen"
150+
instructions: |
151+
Owns optimized inference operations and kernels used by ComfyUI. Check
152+
whether new model math or backend-specific operations duplicate an
153+
existing implementation or belong in comfy-kitchen instead of Core.
143154
learnings:
144155
scope: "auto"

‎comfy/configurable.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
from torch import nn
2+
3+
4+
class ConfigurableModule(nn.Module):
5+
def with_config(self, encoded_config):
6+
"""Return a new module configured from a uint8 JSON tensor, without modifying this module."""
7+
raise NotImplementedError

‎comfy/ldm/lightricks/model.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -327,7 +327,7 @@ def forward(self, x):
327327
# Dropout, so leave it to the stock path whenever it could be active.
328328
if comfy.model_management.in_training:
329329
return self.net(x)
330-
return self.net[2](self.net[0].proj(x), input_act="gelu_tanh")
330+
return comfy.ops.linear_input_act(self.net[2], self.net[0].proj(x), "gelu_tanh")
331331

332332
def apply_rotary_emb(input_tensor, freqs_cis):
333333
rotation_matrix, split_pe = freqs_cis

‎comfy/ldm/minimax/model.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -208,7 +208,7 @@ def __init__(self, hidden, ffn, dtype=None, device=None, operations=None):
208208
self.fc2 = operations.Linear(ffn, hidden, bias=False, dtype=dtype, device=device)
209209

210210
def forward(self, x):
211-
return self.fc2(self.fc1(x), input_act="swiglu")
211+
return comfy.ops.linear_input_act(self.fc2, self.fc1(x), "swiglu")
212212

213213

214214
class AdalnProj(nn.Module):

‎comfy/ldm/minimax/vae.py‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -269,8 +269,8 @@ def __init__(self, dim, mult=4, bias=True, operations=ops):
269269

270270
def forward(self, x, pre_norm, residual, residual_scale):
271271
# norm, gated silu and residual addcmul fold into the INT8 kernels
272-
h = self.w1(x, input_act="rms_norm", act_weight=pre_norm)
273-
return self.w2(h, input_act="swiglu", residual=residual, residual_scale=residual_scale)
272+
h = comfy.ops.linear_input_act(self.w1, x, "rms_norm", act_weight=pre_norm)
273+
return comfy.ops.linear_input_act(self.w2, h, "swiglu", residual=residual, residual_scale=residual_scale)
274274

275275

276276
class Attention(nn.Module):
@@ -289,7 +289,7 @@ def __init__(self, heads, dim_head, bias=True, eps=1e-5, operations=ops):
289289
def forward(self, x, rotary_pos_emb, pre_norm, residual, residual_scale):
290290
batch_size, seq_len, _ = x.shape
291291

292-
qkv = self.to_qkv(x, input_act="rms_norm", act_weight=pre_norm)
292+
qkv = comfy.ops.linear_input_act(self.to_qkv, x, "rms_norm", act_weight=pre_norm)
293293
qkv = qkv.view(batch_size, seq_len, -1, 3 * self.dim_head)
294294
query, key, value = torch.chunk(qkv, 3, dim=-1)
295295

@@ -313,7 +313,7 @@ def forward(self, x, rotary_pos_emb, pre_norm, residual, residual_scale):
313313
out = out.transpose(1, 2).reshape(batch_size, seq_len, -1)
314314
else:
315315
out = optimized_attention(query, key, value, self.heads, skip_reshape=True)
316-
return self.to_out(torch.nan_to_num(out), residual=residual, residual_scale=residual_scale)
316+
return comfy.ops.linear_input_act(self.to_out, torch.nan_to_num(out), None, residual=residual, residual_scale=residual_scale)
317317

318318

319319
class TransformerBlock(nn.Module):

‎comfy/ldm/minimax_music/ar.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@ def __init__(self, hidden_size, intermediate_size, merged_mlp, dtype, device, op
9595

9696
def forward(self, x):
9797
if self.merged_mlp:
98-
return self.down_proj(self.gate_up_proj(x), input_act="swiglu")
98+
return comfy.ops.linear_input_act(self.down_proj, self.gate_up_proj(x), "swiglu")
9999
return self.down_proj(torch.nn.functional.silu(self.gate_proj(x)) * self.up_proj(x))
100100

101101

‎comfy/ldm/modules/attention.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from .sub_quadratic_attention import efficient_dot_product_attention
1818

1919
from comfy import model_management
20+
from comfy.configurable import ConfigurableModule
2021

2122
if model_management.xformers_enabled():
2223
import xformers
@@ -73,12 +74,17 @@ def get_attention_function(name: str, default: Any=...) -> Union[Callable, None]
7374
return REGISTERED_ATTENTION_FUNCTIONS[name]
7475

7576

76-
class ComfyAttention(nn.Module):
77+
class ComfyAttention(ConfigurableModule):
7778
def __init__(self):
7879
super().__init__()
7980
self.config = None
8081
self.function = None
8182

83+
def with_config(self, encoded_config):
84+
attention = ComfyAttention()
85+
attention.load_state_dict({"config": encoded_config})
86+
return attention
87+
8288
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
8389
self.config = None
8490
self.function = None

‎comfy/ldm/qwen_image21/model.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -62,7 +62,7 @@ def __init__(self, dim, hidden_dim, fused=True, dtype=None, device=None, operati
6262

6363
def forward(self, x):
6464
if self.fused:
65-
return self.out(self.gate_up(x), input_act="swiglu")
65+
return comfy.ops.linear_input_act(self.out, self.gate_up(x), "swiglu")
6666
return self.out(F.silu(self.gate_layer(x)) * self.proj(x))
6767

6868

‎comfy/ldm/wan/model.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -196,7 +196,7 @@ class WanFeedForward(nn.Sequential):
196196
"""[Linear, GELU(tanh), Linear], with the GELU folded into the down-projection."""
197197

198198
def forward(self, x):
199-
return self[2](self[0](x), input_act="gelu_tanh")
199+
return comfy.ops.linear_input_act(self[2], self[0](x), "gelu_tanh")
200200

201201

202202
class WanAttentionBlock(nn.Module):

‎comfy/ldm/wan/model_animate2.py‎

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111

1212
import comfy.ldm.common_dit
1313
import comfy.model_management
14+
import comfy.model_prefetch
1415
import comfy.quant_ops
1516
import comfy.utils
1617
from comfy.ldm.flux.math import apply_rope1
@@ -200,16 +201,17 @@ def prefetch(self, i, device, dtype):
200201
stream = None
201202
r = None
202203
if t.device != device:
203-
stream = comfy.model_management.get_offload_stream(device)
204-
cs = comfy.model_management.current_stream(device)
205-
if stream is not None and cs is not None:
206-
# the handed-out stream last waited on the main stream a full rotation ago, which does not cover the previous consumer's reads of this slot; wait now so the copy cannot overwrite a slot still being read
207-
stream.wait_stream(cs)
208204
# two persistent staging buffers per tensor shape instead of a fresh allocation per block (~29 GB of churn per pass at 720p); windows of different lengths get their own pair
209205
buf_key = (tuple(t.shape), cast_dtype if cast_dtype is not None else t.dtype)
210206
if buf_key not in self._staging:
211-
self._staging[buf_key] = [torch.empty(t.shape, dtype=buf_key[1], device=device) for _ in range(2)]
207+
with comfy.model_prefetch.pause_malloc_graph():
208+
self._staging[buf_key] = [torch.empty(t.shape, dtype=buf_key[1], device=device) for _ in range(2)]
212209
r = self._staging[buf_key][i % 2]
210+
stream = comfy.model_management.get_offload_stream(device)
211+
cs = comfy.model_management.current_stream(device)
212+
if stream is not None and cs is not None:
213+
# Wait for staging allocation and the previous consumer's reads before overwriting the buffer.
214+
stream.wait_stream(cs)
213215
self._pending[i] = (comfy.model_management.cast_to(t, cast_dtype, device, non_blocking=True, stream=stream, r=r), stream)
214216

215217
def take(self, i, device, dtype, batch_size):

0 commit comments

Comments
 (0)