Text-to-Image
Diffusers
Safetensors
sdxs-1b / migrate_tapered.py
recoilme's picture
tapered_unet
8587a85
Raw History Blame Contribute Delete
3.98 kB
import os, json, torch
from safetensors.torch import load_file, save_file
from diffusers import UNet2DConditionModel
SRC_DIR = "unet" # trained uniform model (layers=2, trans=[3,3,3,3])
DST_DIR = "unet2" # new tapered model
TAPERED_CFG = {
"block_out_channels": [320, 640, 1280, 1536],
"layers_per_block": [ 4, 3, 2, 2],
"transformer_layers_per_block": [ 4, 3, 2, 1],
"attention_head_dim": [ 10, 10, 10, 12],
}
os.makedirs(DST_DIR, exist_ok=True)
# 1. Load source config, apply tapered overrides
with open(os.path.join(SRC_DIR, "config.json"), "r", encoding="utf-8") as f:
cfg = json.load(f)
for k in ["_class_name", "_diffusers_version", "_name_or_path"]:
cfg.pop(k, None)
cfg.update(TAPERED_CFG)
with open(os.path.join(DST_DIR, "config.json"), "w", encoding="utf-8") as f:
json.dump({
"_class_name": "UNet2DConditionModel",
"_diffusers_version": "0.37.1",
"_name_or_path": "unet2",
**cfg,
}, f, indent=2)
# 2. Load old weights
print("Loading old weights...")
old_state = load_file(os.path.join(SRC_DIR, "diffusion_pytorch_model.safetensors"), device="cpu")
# 3. Create new model
print("Creating tapered model...")
new_model = UNet2DConditionModel(**cfg)
new_state = new_model.state_dict()
old_params = sum(p.numel() for p in old_state.values()) / 1e9
new_params = sum(p.numel() for p in new_state.values()) / 1e9
print(f"Old: {old_params:.3f}B → New: {new_params:.3f}B")
# 4. Copy exact matches
copied, warm = 0, []
for key, old_t in old_state.items():
if key in new_state and old_t.shape == new_state[key].shape:
new_state[key] = old_t.clone()
copied += 1
# 5. Warm-start NEW layers from closest existing layer
# Strategy: for each block, if a resnet.X or transformer.X is new (exists in new but not in old),
# try to copy from X-1. If X-1 also doesn't exist, try X-2, etc.
def find_source(old_state, key, new_state):
"""Try to find a source for a new key by walking backwards through indices."""
# Extract prefix and index: e.g. "down_blocks.0.resnets.2." → ("down_blocks.0.resnets.", 2)
import re
parts = key.rsplit(".", 2)
if len(parts) < 3:
return None
prefix = parts[0] + "." # e.g. "down_blocks.0.resnets."
# Find the numeric index before the last segment
m = re.match(r"(.*\.)(\d+)\.(.+)", key)
if not m:
return None
stem, idx_str, suffix = m.group(1), m.group(2), m.group(3)
idx = int(idx_str)
for src_idx in range(idx - 1, -1, -1):
src_key = f"{stem}{src_idx}.{suffix}"
if src_key in old_state and old_state[src_key].shape == new_state[key].shape:
return src_key
return None
warm_started = []
for key in new_state:
if key not in old_state or old_state[key].shape != new_state[key].shape:
src_key = find_source(old_state, key, new_state)
if src_key:
new_state[key] = old_state[src_key].clone()
warm_started.append((key, src_key))
print(f"Direct copies: {copied}")
print(f"Warm-started: {len(warm_started)} keys")
for k, src in warm_started[:10]:
print(f" {k} ← {src}")
if len(warm_started) > 10:
print(f" ... and {len(warm_started) - 10} more")
# 6. Report stats
total_copied = sum(old_state[k].numel() for k in old_state if k in new_state and old_state[k].shape == new_state[k].shape)
total_ws = sum(new_state[k].numel() for k, _ in warm_started)
total_new = sum(v.numel() for v in new_state.values())
print(f"\nCopied: {total_copied/1e9:.3f}B ({total_copied/total_new*100:.1f}%)")
print(f"Warm-started: {total_ws/1e9:.3f}B ({total_ws/total_new*100:.1f}%)")
print(f"Random init: {(total_new - total_copied - total_ws)/1e9:.3f}B ({(total_new - total_copied - total_ws)/total_new*100:.1f}%)")
# 7. Save
print("\nSaving...")
save_file(new_state, os.path.join(DST_DIR, "diffusion_pytorch_model.safetensors"))
print(f"Done → {DST_DIR}/")