Files

120 lines
5.4 KiB
Python

# Extract the control branch of a trained QwenImage21ControlTransformer2DModel checkpoint.
#
# `train_control.py` / `train_control_distill.py` save the whole transformer (frozen base branch + trainable control
# branch) in the diffusers layout; with FSDP the gathered state dict is written as
# `<checkpoint>/diffusion_pytorch_model.safetensors` (no `transformer/` subdir, no `_control` suffix). The control
# branch is everything that `QwenImage21ControlTransformer2DModel` adds on top of the base model: the
# `control_blocks.*` list (one block per `control_layers` entry) and the `control_img_in.*` input projection. This
# script writes just those tensors to a standalone safetensors file, which can be re-applied onto a fresh base model
# built from `config/qwenimage21/qwenimage21_control.yaml` with `transformer.load_state_dict(..., strict=False)`
# (only the `control_*` keys are consumed; the base keys already match the freshly-loaded base weights).
#
# Usage:
# python scripts/qwenimage21_fun/extract_control_weights.py \
# --model_path /path/to/train_control/checkpoint-xxx/diffusion_pytorch_model.safetensors \
# --output_path /path/to/control_weights.safetensors
import argparse
import json
import os
import torch
from safetensors.torch import load_file, save_file
CONTROL_PREFIXES = ("control_blocks.", "control_img_in.")
# FSDP / DeepSpeed unwrap may leave wrapper prefixes on the keys; strip them to the bare model namespace.
WRAPPER_PREFIXES = ("_fsdp_wrapped_module.", "_fsdp_wrapped_module_", "module.", "_orig_mod.")
def parse_args():
parser = argparse.ArgumentParser(
description="Extract the control-branch weights (control_blocks / control_img_in) of a trained "
"Qwen-Image 2.1 control transformer into a standalone safetensors file."
)
parser.add_argument(
"--model_path", type=str, default="output_dir_qwen_image_21_control/checkpoint-10000/diffusion_pytorch_model.safetensors",
help="Path to the saved transformer: a directory containing diffusion_pytorch_model*.safetensors "
"(e.g. the FSDP `<checkpoint>` dir), or a single .safetensors file.",
)
parser.add_argument(
"--output_path", type=str, default="output_dir_qwen_image_21_control/checkpoint-10000/diffusion_pytorch_model_control.safetensors",
help="Where to write the extracted control weights.",
)
return parser.parse_args()
def resolve_safetensor_files(model_path):
if os.path.isdir(model_path):
shards = sorted(
os.path.join(model_path, name)
for name in os.listdir(model_path)
if name.endswith(".safetensors")
)
if not shards:
raise FileNotFoundError(f"No .safetensors files found under {model_path}.")
return shards
if os.path.isfile(model_path) and model_path.endswith(".safetensors"):
return [model_path]
raise FileNotFoundError(f"--model_path must be a safetensors file or a directory of them, got {model_path}.")
def unwrap_key(key):
changed = True
while changed:
changed = False
for prefix in WRAPPER_PREFIXES:
if key.startswith(prefix):
key = key[len(prefix):]
changed = True
return key
def main():
args = parse_args()
state_dict = {}
for shard in resolve_safetensor_files(args.model_path):
state_dict.update(load_file(shard, device="cpu"))
control_state_dict = {}
for key, value in state_dict.items():
bare_key = unwrap_key(key)
if bare_key.startswith(CONTROL_PREFIXES):
control_state_dict[bare_key] = value.contiguous()
if not control_state_dict:
raise ValueError(
f"No control-branch keys (control_blocks.* / control_img_in.*) found in {args.model_path}; "
"this checkpoint does not look like a Qwen-Image 2.1 control training output."
)
# Carry the branch layout next to the weights so a loader can rebuild the same control model without
# inspecting the full training config. `config.json` of the saved transformer records both fields; keep
# the safetensors metadata strings-only.
metadata = {"format": "pt"}
config_path = os.path.join(args.model_path, "config.json") if os.path.isdir(args.model_path) else None
if config_path is not None and os.path.isfile(config_path):
with open(config_path, "r") as file:
config = json.load(file)
for field in ("control_layers", "control_in_dim"):
if field in config:
metadata[field] = json.dumps(config[field])
os.makedirs(os.path.dirname(os.path.abspath(args.output_path)), exist_ok=True)
save_file(control_state_dict, args.output_path, metadata=metadata)
num_params = sum(value.numel() for value in control_state_dict.values())
block_ids = sorted({
int(key.split(".")[1]) for key in control_state_dict if key.startswith("control_blocks.")
})
print(f"Extracted {len(control_state_dict)} control tensors ({num_params / 1e9:.3f}B params) -> {args.output_path}")
print(f" control_blocks indices: {block_ids}")
if "control_in_dim" in metadata:
print(f" control_in_dim: {metadata['control_in_dim']}, control_layers: {metadata['control_layers']}")
for key in sorted(control_state_dict):
if key.endswith(".weight") and control_state_dict[key].dim() >= 2:
print(f" {key}: {list(control_state_dict[key].shape)} {control_state_dict[key].dtype}")
if __name__ == "__main__":
main()