goldenfox's picture
Marimo Diffusion 0.6B: checkpoint, sampler, OpenAI server, ledger-needle bench
685e018 verified
Raw
History Blame Contribute Delete
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()