Download src/diffusion_lm/sample.py from goldenfox/marimo-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 5.9 kB
-
https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/sample.py
- Command line
-
hf download hf://goldenfox/marimo-diffusion/src/diffusion_lm/sample.py
-
curl -L -o sample.py https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/sample.py
5.9 kB
| """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() | |