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",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
) -> InferenceEngine:
|
) -> InferenceEngine:
|
||||||
if not param_path.exists():
|
if not param_path.exists():
|
||||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||||
@@ -123,6 +124,7 @@ def _create_engine(
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
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}")
|
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||||
return engine
|
return engine
|
||||||
@@ -186,6 +188,7 @@ def run_server(
|
|||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
):
|
):
|
||||||
app = get_app()
|
app = get_app()
|
||||||
app.state.server_config = {
|
app.state.server_config = {
|
||||||
@@ -193,6 +196,7 @@ def run_server(
|
|||||||
"dtype": dtype,
|
"dtype": dtype,
|
||||||
"param_path": param_path,
|
"param_path": param_path,
|
||||||
"max_batch_size": max_batch_size,
|
"max_batch_size": max_batch_size,
|
||||||
|
"max_seq_len": max_seq_len,
|
||||||
}
|
}
|
||||||
uvicorn.run(
|
uvicorn.run(
|
||||||
app,
|
app,
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ class InferenceScheduler:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
cache: Optional[KVCache] = None,
|
cache: Optional[KVCache] = None,
|
||||||
@@ -61,7 +60,6 @@ class InferenceScheduler:
|
|||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=self.max_seq_len,
|
max_seq_len=self.max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self._executor = Executor(
|
self._executor = Executor(
|
||||||
|
|||||||
@@ -135,12 +135,10 @@ class TaskManager:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: int = 8192,
|
max_seq_len: int = 8192,
|
||||||
max_prompt_len: int = 512,
|
|
||||||
):
|
):
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.max_batch_size = max_batch_size
|
self.max_batch_size = max_batch_size
|
||||||
self.max_seq_len = max_seq_len
|
self.max_seq_len = max_seq_len
|
||||||
self.max_prompt_len = max_prompt_len
|
|
||||||
|
|
||||||
self.waiting_queue: Deque[Task] = deque()
|
self.waiting_queue: Deque[Task] = deque()
|
||||||
self.active_tasks: List[Task] = []
|
self.active_tasks: List[Task] = []
|
||||||
@@ -165,10 +163,10 @@ class TaskManager:
|
|||||||
) -> str:
|
) -> str:
|
||||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||||
prompt_ids = self.tokenizer.encode(prompt)
|
prompt_ids = self.tokenizer.encode(prompt)
|
||||||
if len(prompt_ids) > self.max_prompt_len:
|
if len(prompt_ids) > self.max_seq_len:
|
||||||
prompt_ids = prompt_ids[-self.max_prompt_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:
|
if stream_callback:
|
||||||
stream_callback(STOP)
|
stream_callback(STOP)
|
||||||
return task_id
|
return task_id
|
||||||
|
|||||||
@@ -111,8 +111,6 @@ class InferenceEngine:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
|
||||||
page_size: int = 128,
|
|
||||||
cache: Optional[KVCache] = None,
|
cache: Optional[KVCache] = None,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
@@ -122,7 +120,6 @@ class InferenceEngine:
|
|||||||
tokenizer=self.tokenizer,
|
tokenizer=self.tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
|
||||||
cache=cache,
|
cache=cache,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -257,7 +257,6 @@ class TrainContextBuilder:
|
|||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=rollout_batch_size,
|
max_batch_size=rollout_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=max_seq_len or 4096,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
generator = RolloutGenerator(
|
generator = RolloutGenerator(
|
||||||
|
|||||||
@@ -47,7 +47,6 @@ class GenerationBenchmark:
|
|||||||
tokenizer=None,
|
tokenizer=None,
|
||||||
max_batch_size=256,
|
max_batch_size=256,
|
||||||
max_seq_len=config.max_position_embeddings,
|
max_seq_len=config.max_position_embeddings,
|
||||||
max_prompt_len=config.max_position_embeddings,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def run_prefill_benchmark(
|
def run_prefill_benchmark(
|
||||||
|
|||||||
@@ -38,7 +38,6 @@ def processor(
|
|||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=batch_size * num_samples,
|
max_batch_size=batch_size * num_samples,
|
||||||
max_seq_len=cache_len,
|
max_seq_len=cache_len,
|
||||||
max_prompt_len=cache_len,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
print(f"Reading {input_json_file} ...")
|
print(f"Reading {input_json_file} ...")
|
||||||
|
|||||||
+10
-1
@@ -31,7 +31,15 @@ _DTYPES = ["bfloat16", "float16", "float32"]
|
|||||||
default=16,
|
default=16,
|
||||||
help="Maximum batch size for continuous batching.",
|
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)."""
|
"""Launch inference server (OpenAI-compatible API)."""
|
||||||
dtype_map = {
|
dtype_map = {
|
||||||
"bfloat16": torch.bfloat16,
|
"bfloat16": torch.bfloat16,
|
||||||
@@ -51,6 +59,7 @@ def server_command(host, port, reload, param_path, device, dtype, max_batch_size
|
|||||||
dtype=dtype_map[dtype],
|
dtype=dtype_map[dtype],
|
||||||
param_path=Path(param_path),
|
param_path=Path(param_path),
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -402,15 +402,20 @@ def train(
|
|||||||
decay_steps: int,
|
decay_steps: int,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
assert train_type in [
|
if train_type not in [
|
||||||
"seq",
|
"seq",
|
||||||
"sft",
|
"sft",
|
||||||
"dpo",
|
"dpo",
|
||||||
"grpo",
|
"grpo",
|
||||||
"online_grpo",
|
"online_grpo",
|
||||||
"online_dpo",
|
"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":
|
if nprocs > 1 and parallel_mode == "none":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--nprocs > 1 requires --parallel_mode to be 'ddp', 'fsdp', or 'fsdp2'"
|
"--nprocs > 1 requires --parallel_mode to be 'ddp', 'fsdp', or 'fsdp2'"
|
||||||
|
|||||||
@@ -228,7 +228,6 @@ def _make_real_scheduler(device):
|
|||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=8,
|
max_batch_size=8,
|
||||||
max_seq_len=64,
|
max_seq_len=64,
|
||||||
max_prompt_len=64,
|
|
||||||
)
|
)
|
||||||
return scheduler, tokenizer, model
|
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 = TaskManager(tokenizer=t, max_seq_len=16)
|
||||||
tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok))
|
tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok))
|
||||||
assert cb_calls[0] is STOP
|
assert len(cb_calls) == 0
|
||||||
assert len(tm.waiting_queue) == 0
|
assert len(tm.waiting_queue) == 1
|
||||||
|
|
||||||
|
|
||||||
def test_task_manager_remove_task():
|
def test_task_manager_remove_task():
|
||||||
|
|||||||
@@ -109,7 +109,6 @@ def _make_scheduler(model, tokenizer, max_batch_size=8, max_len=128):
|
|||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_len,
|
max_seq_len=max_len,
|
||||||
max_prompt_len=max_len,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user