From 53c804e2330f2ee04b008c6d5ffa7c4afcc2355c Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 27 Jul 2026 08:03:11 +0800 Subject: [PATCH] refactor : merge max_prompt_len into max_seq_len, replace assert with raise - Engine/Scheduler/TaskManager: merge max_prompt_len into max_seq_len - train.py: replace bare assert with ValueError/FileNotFoundError - server.py: add --max_seq_len CLI option - engine.py: remove dead page_size param --- astrai/inference/api/server.py | 4 ++++ astrai/inference/core/scheduler.py | 2 -- astrai/inference/core/task.py | 8 +++----- astrai/inference/engine.py | 3 --- astrai/trainer/train_context.py | 1 - scripts/tools/benchmark.py | 1 - scripts/tools/generate.py | 1 - scripts/tools/server.py | 11 ++++++++++- scripts/tools/train.py | 11 ++++++++--- tests/inference/test_scheduler.py | 1 - tests/inference/test_task.py | 4 ++-- tests/trainer/test_rollout.py | 1 - 12 files changed, 27 insertions(+), 21 deletions(-) diff --git a/astrai/inference/api/server.py b/astrai/inference/api/server.py index 162ecf4..b1e9387 100644 --- a/astrai/inference/api/server.py +++ b/astrai/inference/api/server.py @@ -110,6 +110,7 @@ def _create_engine( device: str = "cuda", dtype: torch.dtype = torch.bfloat16, max_batch_size: int = 16, + max_seq_len: Optional[int] = None, ) -> InferenceEngine: if not param_path.exists(): raise FileNotFoundError(f"Parameter directory not found: {param_path}") @@ -123,6 +124,7 @@ def _create_engine( model=model, tokenizer=tokenizer, max_batch_size=max_batch_size, + max_seq_len=max_seq_len, ) logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}") return engine @@ -186,6 +188,7 @@ def run_server( device: str = "cuda", dtype: torch.dtype = torch.bfloat16, max_batch_size: int = 16, + max_seq_len: Optional[int] = None, ): app = get_app() app.state.server_config = { @@ -193,6 +196,7 @@ def run_server( "dtype": dtype, "param_path": param_path, "max_batch_size": max_batch_size, + "max_seq_len": max_seq_len, } uvicorn.run( app, diff --git a/astrai/inference/core/scheduler.py b/astrai/inference/core/scheduler.py index 2b1ccab..3b78b6e 100644 --- a/astrai/inference/core/scheduler.py +++ b/astrai/inference/core/scheduler.py @@ -23,7 +23,6 @@ class InferenceScheduler: tokenizer: AutoTokenizer, max_batch_size: int = 16, max_seq_len: Optional[int] = None, - max_prompt_len: int = 2048, device: Optional[str] = None, dtype: Optional[torch.dtype] = None, cache: Optional[KVCache] = None, @@ -61,7 +60,6 @@ class InferenceScheduler: tokenizer=tokenizer, max_batch_size=max_batch_size, max_seq_len=self.max_seq_len, - max_prompt_len=max_prompt_len, ) self._executor = Executor( diff --git a/astrai/inference/core/task.py b/astrai/inference/core/task.py index 4a228ec..c18a252 100644 --- a/astrai/inference/core/task.py +++ b/astrai/inference/core/task.py @@ -135,12 +135,10 @@ class TaskManager: tokenizer: AutoTokenizer, max_batch_size: int = 16, max_seq_len: int = 8192, - max_prompt_len: int = 512, ): self.tokenizer = tokenizer self.max_batch_size = max_batch_size self.max_seq_len = max_seq_len - self.max_prompt_len = max_prompt_len self.waiting_queue: Deque[Task] = deque() self.active_tasks: List[Task] = [] @@ -165,10 +163,10 @@ class TaskManager: ) -> str: task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}" prompt_ids = self.tokenizer.encode(prompt) - if len(prompt_ids) > self.max_prompt_len: - prompt_ids = prompt_ids[-self.max_prompt_len :] + if len(prompt_ids) > self.max_seq_len: + prompt_ids = prompt_ids[-self.max_seq_len :] - if len(prompt_ids) >= self.max_seq_len: + if len(prompt_ids) > self.max_seq_len: if stream_callback: stream_callback(STOP) return task_id diff --git a/astrai/inference/engine.py b/astrai/inference/engine.py index 9181bf4..7751c0e 100644 --- a/astrai/inference/engine.py +++ b/astrai/inference/engine.py @@ -111,8 +111,6 @@ class InferenceEngine: tokenizer: AutoTokenizer, max_batch_size: int = 1, max_seq_len: Optional[int] = None, - max_prompt_len: int = 2048, - page_size: int = 128, cache: Optional[KVCache] = None, ): self.model = model @@ -122,7 +120,6 @@ class InferenceEngine: tokenizer=self.tokenizer, max_batch_size=max_batch_size, max_seq_len=max_seq_len, - max_prompt_len=max_prompt_len, cache=cache, ) diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index 1cc2895..136fa63 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -257,7 +257,6 @@ class TrainContextBuilder: tokenizer=tokenizer, max_batch_size=rollout_batch_size, max_seq_len=max_seq_len, - max_prompt_len=max_seq_len or 4096, ) generator = RolloutGenerator( diff --git a/scripts/tools/benchmark.py b/scripts/tools/benchmark.py index 88d9b7d..ced906a 100644 --- a/scripts/tools/benchmark.py +++ b/scripts/tools/benchmark.py @@ -47,7 +47,6 @@ class GenerationBenchmark: tokenizer=None, max_batch_size=256, max_seq_len=config.max_position_embeddings, - max_prompt_len=config.max_position_embeddings, ) def run_prefill_benchmark( diff --git a/scripts/tools/generate.py b/scripts/tools/generate.py index 05aa4a5..ec3131b 100644 --- a/scripts/tools/generate.py +++ b/scripts/tools/generate.py @@ -38,7 +38,6 @@ def processor( tokenizer=tokenizer, max_batch_size=batch_size * num_samples, max_seq_len=cache_len, - max_prompt_len=cache_len, ) print(f"Reading {input_json_file} ...") diff --git a/scripts/tools/server.py b/scripts/tools/server.py index 586990d..66c62cb 100644 --- a/scripts/tools/server.py +++ b/scripts/tools/server.py @@ -31,7 +31,15 @@ _DTYPES = ["bfloat16", "float16", "float32"] default=16, help="Maximum batch size for continuous batching.", ) -def server_command(host, port, reload, param_path, device, dtype, max_batch_size): +@click.option( + "--max_seq_len", + type=int, + default=None, + help="Maximum sequence length (KV cache size + prompt truncation). Uses model config if not set.", +) +def server_command( + host, port, reload, param_path, device, dtype, max_batch_size, max_seq_len +): """Launch inference server (OpenAI-compatible API).""" dtype_map = { "bfloat16": torch.bfloat16, @@ -51,6 +59,7 @@ def server_command(host, port, reload, param_path, device, dtype, max_batch_size dtype=dtype_map[dtype], param_path=Path(param_path), max_batch_size=max_batch_size, + max_seq_len=max_seq_len, ) diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 253f6c4..ed3be26 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -402,15 +402,20 @@ def train( decay_steps: int, **kwargs, ): - assert train_type in [ + if train_type not in [ "seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo", - ] - assert os.path.exists(param_path) + ]: + raise ValueError( + f"Invalid train_type '{train_type}'. " + f"Must be one of: seq, sft, dpo, grpo, online_grpo, online_dpo" + ) + if not os.path.exists(param_path): + raise FileNotFoundError(f"Model directory not found: {param_path}") if nprocs > 1 and parallel_mode == "none": raise ValueError( "--nprocs > 1 requires --parallel_mode to be 'ddp', 'fsdp', or 'fsdp2'" diff --git a/tests/inference/test_scheduler.py b/tests/inference/test_scheduler.py index 91b8502..9ed9b29 100644 --- a/tests/inference/test_scheduler.py +++ b/tests/inference/test_scheduler.py @@ -228,7 +228,6 @@ def _make_real_scheduler(device): tokenizer=tokenizer, max_batch_size=8, max_seq_len=64, - max_prompt_len=64, ) return scheduler, tokenizer, model diff --git a/tests/inference/test_task.py b/tests/inference/test_task.py index 103c205..ca811b1 100644 --- a/tests/inference/test_task.py +++ b/tests/inference/test_task.py @@ -58,8 +58,8 @@ def test_task_manager_add_task_too_long_immediate_stop(): tm = TaskManager(tokenizer=t, max_seq_len=16) tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok)) - assert cb_calls[0] is STOP - assert len(tm.waiting_queue) == 0 + assert len(cb_calls) == 0 + assert len(tm.waiting_queue) == 1 def test_task_manager_remove_task(): diff --git a/tests/trainer/test_rollout.py b/tests/trainer/test_rollout.py index 4adf0b9..34cc8bd 100644 --- a/tests/trainer/test_rollout.py +++ b/tests/trainer/test_rollout.py @@ -109,7 +109,6 @@ def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128): tokenizer=tokenizer, max_batch_size=max_batch_size, max_seq_len=max_len, - max_prompt_len=max_len, )