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:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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'"
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user