| """Command-line smoke evaluator for the shared block diffusion sampler.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| from pathlib import Path |
|
|
| from .core import BlockDiffusionConfig, SAMPLERS, ToyMaskedLMAdapter |
| from .metrics import summarize_result |
|
|
|
|
| def main() -> int: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--method", choices=sorted(SAMPLERS), default="confidence") |
| parser.add_argument("--steps", type=int, default=8) |
| parser.add_argument("--block-size", type=int, default=16) |
| parser.add_argument("--num-blocks", type=int, default=1) |
| parser.add_argument("--remask-ratio", type=float, default=0.5) |
| parser.add_argument("--use-cache", action="store_true") |
| parser.add_argument("--output", type=Path) |
| args = parser.parse_args() |
|
|
| config = BlockDiffusionConfig( |
| block_size=args.block_size, |
| num_blocks=args.num_blocks, |
| steps=args.steps, |
| remask_ratio=args.remask_ratio, |
| use_cache=args.use_cache, |
| ) |
| adapter = ToyMaskedLMAdapter(vocab_size=config.vocab_size, mask_token_id=config.mask_token_id) |
| sampler = SAMPLERS[args.method](adapter, config) |
| result = sampler.decode() |
| summary = summarize_result(result) |
| text = json.dumps(summary, indent=2, ensure_ascii=False) |
| if args.output: |
| args.output.parent.mkdir(parents=True, exist_ok=True) |
| args.output.write_text(text + "\n", encoding="utf-8") |
| print(text) |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|
|
|