| 424 | |
| 425 | |
| 426 | def build_parser() -> argparse.ArgumentParser: |
| 427 | parser = argparse.ArgumentParser(description=__doc__) |
| 428 | parser.add_argument( |
| 429 | "--op-source", |
| 430 | choices=["registry", "native", "triton", "sm90"], |
| 431 | default="registry", |
| 432 | ) |
| 433 | parser.add_argument("--dtype", default="bf16", help="bf16, fp16, or fp32") |
| 434 | parser.add_argument("--reference-mode", choices=["matching", "fp32"], default="matching") |
| 435 | parser.add_argument("--seed", type=int, default=1234) |
| 436 | parser.add_argument("--tokens", type=int, default=128) |
| 437 | parser.add_argument("--hidden-size", type=int, default=256) |
| 438 | parser.add_argument("--vocab-size", type=int, default=4096) |
| 439 | parser.add_argument("--no-bias", action="store_true") |
| 440 | parser.add_argument("--uneven-shards", action="store_true") |
| 441 | parser.add_argument("--atol", type=float, default=None) |
| 442 | parser.add_argument("--rtol", type=float, default=None) |
| 443 | parser.add_argument("--run-stress", action="store_true") |
| 444 | parser.add_argument("--stress-tokens", type=int, default=4096) |
| 445 | parser.add_argument("--stress-hidden-size", type=int, default=2048) |
| 446 | parser.add_argument("--stress-vocab-size", type=int, default=32768) |
| 447 | return parser |
| 448 | |
| 449 | |
| 450 | def main() -> None: |