| 23 | |
| 24 | |
| 25 | def get_args(): |
| 26 | parser = argparse.ArgumentParser(description="CPU bitsandbytes optimizer training") |
| 27 | parser.add_argument("--model", type=str, default="JackFram/llama-68m") |
| 28 | parser.add_argument("--dataset", type=str, default="yahma/alpaca-cleaned") |
| 29 | parser.add_argument( |
| 30 | "--optimizer", |
| 31 | type=str, |
| 32 | default="adamw", |
| 33 | choices=[ |
| 34 | "adamw", |
| 35 | "adamw8bit", |
| 36 | "adamw32bit", |
| 37 | "adam", |
| 38 | "adam8bit", |
| 39 | "adam32bit", |
| 40 | "sgd", |
| 41 | "sgd8bit", |
| 42 | "lion", |
| 43 | "lion8bit", |
| 44 | "rmsprop", |
| 45 | "rmsprop8bit", |
| 46 | "adagrad", |
| 47 | "adagrad8bit", |
| 48 | "lamb", |
| 49 | "lars", |
| 50 | ], |
| 51 | ) |
| 52 | parser.add_argument("--lr", type=float, default=2e-4) |
| 53 | parser.add_argument("--batch_size", type=int, default=2) |
| 54 | parser.add_argument("--max_length", type=int, default=128) |
| 55 | parser.add_argument("--steps", type=int, default=30) |
| 56 | parser.add_argument("--log_interval", type=int, default=5) |
| 57 | parser.add_argument("--compare", action="store_true", help="Compare bnb AdamW vs torch AdamW") |
| 58 | parser.add_argument("--use_trainer", action="store_true", help="Use HF Trainer instead of manual training loop") |
| 59 | parser.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp32"]) |
| 60 | return parser.parse_args() |
| 61 | |
| 62 | |
| 63 | def format_alpaca(example): |