从命令行解析参数
()
| 45 | |
| 46 | |
| 47 | def parse_args(): |
| 48 | """ |
| 49 | 从命令行解析参数 |
| 50 | """ |
| 51 | parser = argparse.ArgumentParser("Cache Messager") |
| 52 | parser.add_argument( |
| 53 | "--splitwise_role", |
| 54 | type=str, |
| 55 | default="mixed", |
| 56 | help="splitwise role, can be decode, prefill or mixed", |
| 57 | ) |
| 58 | parser.add_argument("--rank", type=int, default=0, help="local tp rank id") |
| 59 | parser.add_argument("--device_id", type=int, default=0, help="device id") |
| 60 | parser.add_argument("--num_layers", type=int, default=1, help="model num layers") |
| 61 | parser.add_argument("--key_cache_shape", type=str, default="", help="key cache shape") |
| 62 | parser.add_argument("--value_cache_shape", type=str, default="", help="value cache shape") |
| 63 | parser.add_argument("--rdma_port", type=str, default="", help="rmda port") |
| 64 | parser.add_argument("--mp_num", type=int, default=1, help="number of model parallel, i.e. tp_size, tp_num") |
| 65 | parser.add_argument("--ipc_suffix", type=str, default=None, help="ipc suffix") |
| 66 | parser.add_argument( |
| 67 | "--protocol", |
| 68 | type=str, |
| 69 | default="ipc", |
| 70 | help="cache transfer protocol, only surport ipc now", |
| 71 | ) |
| 72 | parser.add_argument("--pod_ip", type=str, default="0.0.0.0", help="pod ip") |
| 73 | parser.add_argument("--cache_queue_port", type=int, default=9924, help="cache queue port") |
| 74 | parser.add_argument( |
| 75 | "--engine_worker_queue_port", |
| 76 | type=int, |
| 77 | default=9923, |
| 78 | help="engine worker queue port", |
| 79 | ) |
| 80 | parser.add_argument( |
| 81 | "--cache_dtype", |
| 82 | type=str, |
| 83 | default="bfloat16", |
| 84 | choices=["uint8", "bfloat16", "block_wise_fp8"], |
| 85 | help="cache dtype", |
| 86 | ) |
| 87 | parser.add_argument( |
| 88 | "--default_dtype", |
| 89 | type=str, |
| 90 | default="bfloat16", |
| 91 | choices=["float16", "bfloat16", "uint8", "int8"], |
| 92 | help="paddle default dtype, cache manager only support float16、bfloat16、int8 and uint8 now", |
| 93 | ) |
| 94 | parser.add_argument( |
| 95 | "--speculative_config", |
| 96 | type=json.loads, |
| 97 | default="{}", |
| 98 | help="speculative config", |
| 99 | ) |
| 100 | parser.add_argument("--local_data_parallel_id", type=int, default=0) |
| 101 | |
| 102 | args = parser.parse_args() |
| 103 | return args |
| 104 |
no test coverage detected