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}/")