24 lines
1.1 KiB
Python
24 lines
1.1 KiB
Python
"""Capture block-0 AdaLN gates for direct residual parity diagnostics."""
|
|
|
|
from pathlib import Path
|
|
|
|
|
|
model = Path("/opt/ComfyUI/comfy/ldm/minimax/model.py")
|
|
source = model.read_text(encoding="utf-8")
|
|
old = (
|
|
'torch.save({"norm": norm1.detach().cpu(), "shift": shift_msa.detach().cpu(), '
|
|
'"scale": scale_msa.detach().cpu(), "effective_weight": effective_weight.detach().cpu(), '
|
|
'"effective_bias": effective_bias.detach().cpu()}, os.path.join(capture_dir, "block0_norm1_adaln.pt"))'
|
|
)
|
|
new = (
|
|
'torch.save({"norm": norm1.detach().cpu(), "shift": shift_msa.detach().cpu(), '
|
|
'"scale": scale_msa.detach().cpu(), "gate_msa": gate_msa.detach().cpu(), '
|
|
'"gate_mlp": gate_mlp.detach().cpu(), "effective_weight": effective_weight.detach().cpu(), '
|
|
'"effective_bias": effective_bias.detach().cpu()}, os.path.join(capture_dir, "block0_norm1_adaln.pt"))'
|
|
)
|
|
if source.count(old) == 1:
|
|
source = source.replace(old, new)
|
|
elif new not in source:
|
|
raise RuntimeError("Unable to locate block-0 AdaLN capture payload.")
|
|
model.write_text(source, encoding="utf-8")
|
|
print("Added H3 block-0 AdaLN gate capture.")
|