| 36 | |
| 37 | # -------------------------- ARGPARSE / CONFIG ------------------------ # |
| 38 | def parse_args() -> argparse.Namespace: |
| 39 | p = argparse.ArgumentParser() |
| 40 | p.add_argument("--exp_name", default="rank_iter0") |
| 41 | p.add_argument("--dataset", default="data/synthetic_data/train/iter0_train.json") |
| 42 | p.add_argument("--output_dir", default="results/query_server") |
| 43 | p.add_argument("--server_host", default="127.0.0.1") |
| 44 | p.add_argument("--zmq_port", type=int, default=5555) |
| 45 | |
| 46 | p.add_argument("--n_articles", type=int, default=3) |
| 47 | p.add_argument("--k_completions", type=int, default=5) |
| 48 | p.add_argument("--eval_times", type=int, default=3) |
| 49 | |
| 50 | # LoRA / optimization hyperparams |
| 51 | p.add_argument("--lora_rank", type=int, default=32) |
| 52 | p.add_argument("--lora_alpha", type=int, default=64) |
| 53 | p.add_argument("--lora_dropout", type=float, default=0) |
| 54 | p.add_argument("--finetune_epochs", type=int, default=10) |
| 55 | p.add_argument("--finetune_lr", type=float, default=1e-3) |
| 56 | p.add_argument("--batch_size", type=int, default=1) |
| 57 | p.add_argument("--gradient_accumulation_steps", type=int, default=1) |
| 58 | p.add_argument("--end_mask_substring", default="") |
| 59 | p.add_argument("--split_newlines", action="store_true") |
| 60 | p.add_argument("--chain_of_thought", action="store_true") |
| 61 | p.add_argument("--reward_mode", choices=["ttt", "proxy", "both"], default="ttt") |
| 62 | return p.parse_args() |
| 63 | |
| 64 | |
| 65 | def send_round_trip( |