| """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: |
| |
| |
| checkpoint = torch.load( |
| checkpoint_path, |
| map_location="cpu", |
| weights_only=False, |
| mmap=True, |
| ) |
| except TypeError: |
| 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": |
| |
| |
| 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, |
| 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() |
|
|