Text Generation
Transformers
Safetensors
English
metadiffusion
diffusion
diffusion-lm
ar-to-diffusion
custom_code
CodeSoft's picture
Upload 9 files
d6f5237 verified
Raw
History Blame Contribute Delete
7.69 kB
#!/usr/bin/env python3
"""export_hf.py: export a MetaDiffusion-600M checkpoint to a release dir.
Output:
model.safetensors fp16, "model."-prefixed keys
config.json no dtype key; vocab fields = weight rows;
auto_map -> hf_modeling.py custom classes
generation_config.json denoising defaults
tokenizer/ Qwen3 tokenizer + [MASK] + rainbow tokens
hf_modeling.py standalone modeling (AutoModelForCausalLM,
GenerationMixin with iterative denoising)
scripts/ self-contained pipeline copy
Usage:
python export_hf.py --checkpoint checkpoints/step_30000.pt \
--tokenizer data/tokenizer --output MetaDiffusion-600M-Instruct-v1
Then load with:
AutoModelForCausalLM.from_pretrained(dir, trust_remote_code=True)
"""
import argparse
import json
import shutil
import sys
from pathlib import Path
import torch
from safetensors.torch import save_file
from transformers import AutoTokenizer
sys.path.insert(0, str(Path(__file__).resolve().parent))
from model import MetaDiffusionConfig # noqa: E402
GENERATION_CONFIG = {
"temperature": 0.7,
"repetition_penalty": 1.5,
"num_steps": 128,
"max_new_tokens": 96,
"top_p": 0.0, # truncation sampling: min-p is the release default
"min_p": 0.1,
"im_end_bias": 2.0,
"im_end_bias_t": 0.3,
"do_sample": True,
"transformers_version": "4.49.0",
}
PLAIN_CHAT_TEMPLATE = (
"{% for message in messages %}{{ '<|im_start|>' + message['role'] }}\n"
"{{ message['content'] }}<|im_end|>\n"
"{% endfor %}{% if add_generation_prompt %}<|im_start|>assistant\n"
"{% endif %}"
)
def remap_state_dict(state_dict, dtype="bf16"):
cast = {"bf16": torch.bfloat16, "fp16": torch.float16,
"fp32": torch.float32}[dtype]
new_dict = {}
for key, tensor in state_dict.items():
key = key.replace("_orig_mod.", "", 1) if key.startswith("_orig_mod.") else key
new_dict["model." + key] = tensor.to(cast)
return new_dict
def package_scripts(out):
src = Path(__file__).resolve().parent
scripts_dir = out / "scripts"
scripts_dir.mkdir(parents=True, exist_ok=True)
for name in ["model.py", "convert.py", "prepare_data.py", "train.py",
"chat.py", "eval.py", "export_hf.py", "hf_modeling.py"]:
cand = src / name
if cand.exists():
shutil.copy2(cand, scripts_dir / name)
req = scripts_dir / "requirements.txt"
if not req.exists():
req.write_text("torch>=2.2\ntransformers>=4.49\nsafetensors>=0.4\n"
"datasets>=2.18\nnumpy>=1.26\n")
print(f"[*] Packaged scripts -> {scripts_dir}")
def export(checkpoint_path, tokenizer_dir, output_dir, dtype="bf16"):
out = Path(output_dir)
out.mkdir(parents=True, exist_ok=True)
print(f"[*] Loading checkpoint {checkpoint_path}")
ckpt = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
config = MetaDiffusionConfig(
**{k: v for k, v in ckpt["config"].items()
if k in MetaDiffusionConfig.__dataclass_fields__})
print("[*] Remapping state dict...")
state_dict = remap_state_dict(ckpt["model_state_dict"], dtype=dtype)
save_file(state_dict, out / "model.safetensors")
print(f"[*] Saved {len(state_dict)} tensors -> {out / 'model.safetensors'} "
f"({dtype})")
# Vocab fields must match the actual weight rows
n_vocab = state_dict["model.lm_head.weight"].shape[0]
config.vocab_size = n_vocab
config.mask_vocab_size = n_vocab
print(f"[*] Vocab in config: {n_vocab} (matches weights)")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_dir)
eos_ids = [tokenizer.eos_token_id] if tokenizer.eos_token_id is not None else []
im_end = tokenizer.convert_tokens_to_ids("<|im_end|>")
if im_end != tokenizer.unk_token_id and im_end not in eos_ids:
eos_ids.append(im_end)
eos_ids = [e for e in eos_ids if e is not None]
print(f"[*] eos ids: {eos_ids}")
# partial-byte vocab entries that cannot decode to valid UTF-8 (the
# literal replacement-char garbage): hard-banned at generation
bad_ids = []
for i in range(len(tokenizer)):
s = tokenizer.decode([i], skip_special_tokens=True)
if s and all(c == "\uFFFD" for c in s):
bad_ids.append(i)
config_dict = config.__dict__.copy()
config_dict.pop("dtype", None) # transformers chokes on "torch.float32" strings
config_dict["model_type"] = "metadiffusion"
config_dict["architectures"] = ["MetaDiffusion600MForCausalLM"]
config_dict["auto_map"] = {
"AutoConfig": "hf_modeling.MetaDiffusion600MConfig",
"AutoModelForCausalLM": "hf_modeling.MetaDiffusion600MForCausalLM",
}
config_dict["eos_token_id"] = eos_ids
config_dict["rainbow_token_ids"] = list(range(config.mask_token_id + 1,
config.mask_token_id + 8))
config_dict["invalid_utf8_token_ids"] = bad_ids
with open(out / "config.json", "w") as f:
json.dump(config_dict, f, indent=2)
print(f"[*] Saved config.json (mask_token_id={config.mask_token_id}, "
f"{len(bad_ids)} banned invalid-UTF8 tokens)")
gen_config = dict(GENERATION_CONFIG)
gen_config["eos_token_id"] = eos_ids
gen_config["pad_token_id"] = config.pad_token_id
gen_config["mask_token_id"] = config.mask_token_id
with open(out / "generation_config.json", "w") as f:
json.dump(gen_config, f, indent=2)
shutil.copytree(tokenizer_dir, out / "tokenizer", dirs_exist_ok=True)
print(f"[*] Copied tokenizer -> {out / 'tokenizer'}")
# pin the plain chat template in the exported tokenizer (the saved one is
# empty, which silently falls back to the think-injecting Qwen3 default)
tok_cfg_path = out / "tokenizer" / "tokenizer_config.json"
if tok_cfg_path.exists():
tc = json.loads(tok_cfg_path.read_text())
tc["chat_template"] = PLAIN_CHAT_TEMPLATE
tok_cfg_path.write_text(json.dumps(tc, indent=2, ensure_ascii=False))
print("[*] Pinned plain chat_template in exported tokenizer")
# chat_template.jinja takes precedence over the config string since
# transformers 5.x; overwrite it so both sources carry the plain
# template (the Qwen3 default jinja injects <think> blocks).
(out / "tokenizer" / "chat_template.jinja").write_text(
PLAIN_CHAT_TEMPLATE)
print("[*] Pinned plain chat_template.jinja")
src = Path(__file__).resolve().parent
shutil.copy2(src / "hf_modeling.py", out / "hf_modeling.py")
print("[*] Copied hf_modeling.py (trust_remote_code)")
package_scripts(out)
print(f"[*] Done: {out}")
print(" Load with: AutoModelForCausalLM.from_pretrained("
f"'{out}', trust_remote_code=True)")
print(" (Write the README yourself; export never touches it.)")
def main():
p = argparse.ArgumentParser(description="Export MetaDiffusion-600M release dir")
p.add_argument("--checkpoint", required=True)
p.add_argument("--tokenizer", default="data/tokenizer")
p.add_argument("--output", required=True)
p.add_argument("--dtype", default="bf16", choices=["bf16", "fp16", "fp32"],
help="Weight dtype (default bf16: matches training dtype and "
"survives timestep extrapolation; fp16 overflows to NaN "
"on curriculum-trained checkpoints)")
args = p.parse_args()
export(args.checkpoint, args.tokenizer, args.output, dtype=args.dtype)
if __name__ == "__main__":
main()