Spaces:
Running on Zero
Running on Zero
File size: 8,231 Bytes
2dc3625 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 | """Install SCoPE attention in both Wan2.2-A14B experts."""
import torch
from scope.modeling import SelfAttentionWithSCoPE, create_scope_block_forward
def _patch_expert(
dit,
method: str,
height: int,
width: int,
copy_self_attn_weights: bool,
plucker_init: str,
plucker_init_scale: float,
plucker_mlp_hidden: int,
plucker_scale: float,
gate_init_bias: float,
disable_spatial_rope: bool,
cam_residual_layers: list[int] | None,
scale_gate_hidden: int,
log_scale_aug_prob: float,
log_scale_aug_range: tuple,
log_prefix: str = "[SCoPE]",
):
"""Replace self-attention in one DiT expert."""
dit.camera_condition = method
num_blocks = len(dit.blocks)
if cam_residual_layers is None:
cam_residual_set = set(range(num_blocks))
layer_desc = "all"
else:
cam_residual_set = set(cam_residual_layers)
layer_desc = str(sorted(cam_residual_set))
print(
f"{log_prefix} method={method}, plucker_init={plucker_init}, "
f"plucker_init_scale={plucker_init_scale}, "
f"plucker_mlp_hidden={plucker_mlp_hidden}, plucker_scale={plucker_scale}, "
f"gate_init_bias={gate_init_bias}, scale_gate_hidden={scale_gate_hidden}, "
f"log_scale_aug_prob={log_scale_aug_prob}, "
f"log_scale_aug_range={log_scale_aug_range}, "
f"disable_spatial_rope={disable_spatial_rope}, "
f"cam_residual_layers={layer_desc} ({len(cam_residual_set)}/{num_blocks} blocks)"
)
for i, block in enumerate(dit.blocks):
original_attn = block.self_attn
enable_cam_residual = i in cam_residual_set
new_attn = SelfAttentionWithSCoPE(
dim=dit.dim,
num_heads=block.num_heads,
eps=1e-6,
plucker_init=plucker_init,
plucker_init_scale=plucker_init_scale,
plucker_mlp_hidden=plucker_mlp_hidden,
plucker_scale=plucker_scale,
gate_init_bias=gate_init_bias,
disable_spatial_rope=disable_spatial_rope,
enable_cam_residual=enable_cam_residual,
scale_gate_hidden=scale_gate_hidden,
log_scale_aug_prob=log_scale_aug_prob,
log_scale_aug_range=log_scale_aug_range,
)
if copy_self_attn_weights:
new_attn.q.weight.data = original_attn.q.weight.data.clone()
if original_attn.q.bias is not None and new_attn.q.bias is not None:
new_attn.q.bias.data = original_attn.q.bias.data.clone()
new_attn.k.weight.data = original_attn.k.weight.data.clone()
if original_attn.k.bias is not None and new_attn.k.bias is not None:
new_attn.k.bias.data = original_attn.k.bias.data.clone()
new_attn.v.weight.data = original_attn.v.weight.data.clone()
if original_attn.v.bias is not None and new_attn.v.bias is not None:
new_attn.v.bias.data = original_attn.v.bias.data.clone()
new_attn.o.weight.data = original_attn.o.weight.data.clone()
if original_attn.o.bias is not None and new_attn.o.bias is not None:
new_attn.o.bias.data = original_attn.o.bias.data.clone()
new_attn.norm_q.weight.data = original_attn.norm_q.weight.data.clone()
new_attn.norm_k.weight.data = original_attn.norm_k.weight.data.clone()
block.self_attn = new_attn
forward_fn = create_scope_block_forward()
for block in dit.blocks:
block.forward = forward_fn.__get__(block, block.__class__)
def patch_scope(
pipe,
method,
height,
width,
copy_self_attn_weights: bool = True,
plucker_init: str = "zero",
plucker_init_scale: float = 0.01,
plucker_mlp_hidden: int = 0,
plucker_scale: float = 0.0,
gate_init_bias: float = -2.0,
disable_spatial_rope: bool = False,
cam_residual_layers: list[int] | None = None,
scale_gate_hidden: int = 0,
log_scale_aug_prob: float = 0.0,
log_scale_aug_range: tuple = (-1.2, 1.6),
**kwargs,
):
"""Patch the high- and low-noise experts and return trainable key patterns."""
if getattr(pipe, "dit2", None) is None:
raise RuntimeError("SCoPE requires both Wan2.2-A14B experts; low-noise expert is missing.")
common_kwargs = {
"method": method,
"height": height,
"width": width,
"copy_self_attn_weights": copy_self_attn_weights,
"plucker_init": plucker_init,
"plucker_init_scale": plucker_init_scale,
"plucker_mlp_hidden": plucker_mlp_hidden,
"plucker_scale": plucker_scale,
"gate_init_bias": gate_init_bias,
"disable_spatial_rope": disable_spatial_rope,
"cam_residual_layers": cam_residual_layers,
"scale_gate_hidden": scale_gate_hidden,
"log_scale_aug_prob": log_scale_aug_prob,
"log_scale_aug_range": log_scale_aug_range,
}
_patch_expert(pipe.dit, log_prefix="[SCoPE/dit]", **common_kwargs)
_patch_expert(pipe.dit2, log_prefix="[SCoPE/dit2]", **common_kwargs)
keywords = ["plucker_pe", "self_attn", "norm3", "ffn"]
return keywords
def validate_official_low_expert(dit) -> None:
"""Verify that the patched low-noise expert remains an exact zero-delta model."""
for index, block in enumerate(dit.blocks):
positional_encoding = block.self_attn.plucker_pe
if positional_encoding.enable_cam_residual:
raise RuntimeError(f"Low expert block {index} has camera residual enabled")
q_output = (
positional_encoding.eq[2] if positional_encoding.use_mlp else positional_encoding.eq
)
k_output = (
positional_encoding.ek[2] if positional_encoding.use_mlp else positional_encoding.ek
)
if torch.count_nonzero(q_output.weight).item() != 0:
raise RuntimeError(f"Low expert block {index} has non-zero SCoPE query weights")
if torch.count_nonzero(k_output.weight).item() != 0:
raise RuntimeError(f"Low expert block {index} has non-zero SCoPE key weights")
def enable_scope_grad(pipe, keywords, expert: str = "high_noise_model"):
"""Enable gradients for the selected expert; the released run trained high noise only."""
pipe.eval()
pipe.requires_grad_(False)
if getattr(pipe, "dit2", None) is None:
raise RuntimeError("SCoPE training requires both Wan2.2-A14B experts.")
if expert == "high_noise_model":
targets = [(pipe.dit, "")]
elif expert == "low_noise_model":
targets = [(pipe.dit2, "[dit2] ")]
elif expert == "both":
targets = [(pipe.dit, ""), (pipe.dit2, "[dit2] ")]
else:
raise ValueError(f"Unknown expert selection: {expert}")
if keywords == "*":
for dit, _ in targets:
dit.train()
dit.requires_grad_(True)
else:
for dit, prefix in targets:
for name, module in dit.named_modules():
if any(keyword in name for keyword in keywords):
print(f"Trainable: {prefix}{name}")
module.train()
module.requires_grad_(True)
# DiTBlock.modulation / Head.modulation 是直接挂在 block/head 上的
# nn.Parameter (非子 Module), named_modules 匹配不到, 需按 param 名补开.
for name, param in dit.named_parameters():
if not param.requires_grad and any(keyword in name for keyword in keywords):
print(f"Trainable param: {prefix}{name}")
param.requires_grad = True
trainable_params = 0
seen_params = set()
for dit, _ in targets:
for _, module in dit.named_modules():
for param in module.parameters():
if param.requires_grad and id(param) not in seen_params:
trainable_params += param.numel()
seen_params.add(id(param))
print(f"Total number of trainable parameters (dit+dit2): {trainable_params:,}")
|