| 20 | |
| 21 | |
| 22 | def get_args(): |
| 23 | parser = argparse.ArgumentParser(description="Paged Optimizer Memory Benchmark") |
| 24 | parser.add_argument("--hidden_size", type=int, default=1024) |
| 25 | parser.add_argument("--num_layers", type=int, default=12) |
| 26 | parser.add_argument("--intermediate_size", type=int, default=2752) |
| 27 | parser.add_argument("--num_heads", type=int, default=16) |
| 28 | parser.add_argument("--vocab_size", type=int, default=32000) |
| 29 | parser.add_argument("--seq_len", type=int, default=128) |
| 30 | parser.add_argument("--batch_size", type=int, default=2) |
| 31 | parser.add_argument("--train_steps", type=int, default=5) |
| 32 | parser.add_argument("--device", type=str, default="xpu") |
| 33 | parser.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16", "fp32"]) |
| 34 | return parser.parse_args() |
| 35 | |
| 36 | |
| 37 | def get_torch_dtype(name): |