"""Generate fixed-length text by iteratively unmasking tokens.""" from __future__ import annotations import argparse import gc import hashlib from pathlib import Path import torch from diffusion_lm.config import ModelConfig from diffusion_lm.diffusion import iterative_unmask from diffusion_lm.model import DiffusionTransformer from diffusion_lm.tokenizer import load_tokenizer, special_token_id, special_token_ids from diffusion_lm.train import resolve_device def load_model(checkpoint_path: str | Path, device: torch.device) -> DiffusionTransformer: try: # mmap keeps unused optimizer tensors in a full training checkpoint off # resident RAM. The compact inference export remains the preferred input. checkpoint = torch.load( checkpoint_path, map_location="cpu", weights_only=False, mmap=True, ) except TypeError: # PyTorch versions before mmap= support. checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False) checkpoint_format = checkpoint.get("format") if checkpoint_format not in { "mini-diffusion-lm-checkpoint-v1", "mini-diffusion-lm-inference-v1", }: raise ValueError("unsupported checkpoint format") model_config = ModelConfig(**checkpoint["config"]["model"]) model = DiffusionTransformer(model_config) if checkpoint_format == "mini-diffusion-lm-inference-v1" and device.type != "cpu": # The weights-only export is BF16. Keep that dtype on accelerators instead # of silently expanding a 1B model back to FP32 during load_state_dict. first_weight = next(iter(checkpoint["model"].values())) if first_weight.is_floating_point(): model = model.to(device=device, dtype=first_weight.dtype) model.load_state_dict(checkpoint["model"]) model.tokenizer_sha256 = checkpoint.get("tokenizer_sha256") del checkpoint gc.collect() return model.to(device).eval() def generate( model: DiffusionTransformer, tokenizer_path: str | Path, *, prompt: str = "", generation_length: int = 128, num_samples: int = 1, steps: int = 64, temperature: float = 1.0, strategy: str = "ancestral", seed: int = 1337, ) -> list[str]: tokenizer = load_tokenizer(tokenizer_path) tokenizer_hash = hashlib.sha256(Path(tokenizer_path).read_bytes()).hexdigest() if model.tokenizer_sha256 is not None and tokenizer_hash != model.tokenizer_sha256: raise ValueError("tokenizer file does not match the tokenizer used for training") if tokenizer.get_vocab_size(with_added_tokens=True) != model.config.vocab_size: raise ValueError("tokenizer vocabulary does not match the checkpoint") mask_id = special_token_id(tokenizer, "mask") if mask_id != model.config.mask_token_id: raise ValueError("tokenizer mask id does not match the checkpoint") if generation_length <= 0 or num_samples <= 0: raise ValueError("generation_length and num_samples must be positive") prompt_ids = tokenizer.encode(prompt).ids if prompt else [] total_length = len(prompt_ids) + generation_length if total_length > model.config.max_seq_len: raise ValueError( f"prompt plus generation uses {total_length} tokens, but model limit is " f"{model.config.max_seq_len}" ) device = next(model.parameters()).device input_ids = torch.full( (num_samples, total_length), model.config.mask_token_id, dtype=torch.long, device=device, ) if prompt_ids: input_ids[:, : len(prompt_ids)] = torch.tensor(prompt_ids, device=device) torch.manual_seed(seed) if device.type == "cuda": torch.cuda.manual_seed_all(seed) elif device.type == "mps" and hasattr(torch.mps, "manual_seed"): torch.mps.manual_seed(seed) role_ids = special_token_ids(tokenizer) blocked = tuple(role_ids[role] for role in ("pad", "unk", "bos", "mask")) result = iterative_unmask( model, input_ids, model.config.mask_token_id, steps=steps, temperature=temperature, strategy=strategy, # type: ignore[arg-type] blocked_token_ids=blocked, ).cpu() eos_id = special_token_id(tokenizer, "eos") texts: list[str] = [] for row in result.tolist(): if eos_id in row[len(prompt_ids) :]: eos_position = row.index(eos_id, len(prompt_ids)) row = row[:eos_position] texts.append(tokenizer.decode(row, skip_special_tokens=True)) return texts def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--tokenizer", type=Path, required=True) parser.add_argument("--prompt", default="") parser.add_argument("--length", type=int, default=128, help="number of completion tokens") parser.add_argument("--num-samples", type=int, default=1) parser.add_argument("--steps", type=int, default=64) parser.add_argument("--temperature", type=float, default=1.0) parser.add_argument("--strategy", choices=("ancestral", "confidence"), default="ancestral") parser.add_argument("--seed", type=int, default=1337) parser.add_argument("--device", default="auto") args = parser.parse_args() device = resolve_device(args.device) model = load_model(args.checkpoint, device) texts = generate( model, args.tokenizer, prompt=args.prompt, generation_length=args.length, num_samples=args.num_samples, steps=args.steps, temperature=args.temperature, strategy=args.strategy, seed=args.seed, ) for index, text in enumerate(texts, start=1): print(f"[{index}] {text}") if __name__ == "__main__": main()