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
This commit is contained in:
2026-07-27 08:05:11 +08:00
parent 05c7432964
commit 53c804e233
12 changed files with 27 additions and 21 deletions
+4
View File
@@ -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,
-2
View File
@@ -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(
+3 -5
View File
@@ -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
-3
View File
@@ -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,
)
-1
View File
@@ -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(
-1
View File
@@ -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(
-1
View File
@@ -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} ...")
+10 -1
View File
@@ -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,
)
+8 -3
View File
@@ -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'"
-1
View File
@@ -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
+2 -2
View File
@@ -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():
-1
View File
@@ -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,
)