refactor: 设计模式优化 inference 模块导入结构
- 新建 cache.py:SlotAllocator 对象池 + PrefixCacheManager - 新建 sampling.py:Temperature/TopK/TopP 可组合策略 - TaskStatus 改用 Enum,GenerationParams 值对象模式 - _STOP 移至 cache.py,解除 engine→scheduler 轻量耦合 - 更新测试导入路径,ruff 格式检查通过
This commit is contained in:
+49
-12
@@ -1,22 +1,41 @@
|
||||
"""Unified inference engine for continuous batching."""
|
||||
"""Unified inference engine for continuous batching.
|
||||
|
||||
Layers:
|
||||
- GenerationParams: Immutable value object for sampling parameters.
|
||||
- GenerationRequest: User-facing request DTO with validation.
|
||||
- _Result: Thread-safe token accumulator (Observer pattern).
|
||||
- InferenceEngine: Facade over InferenceScheduler + async wrapper.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.inference.scheduler import _STOP, InferenceScheduler
|
||||
from astrai.inference.cache import _STOP
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GenerationParams:
|
||||
"""Immutable value object for sampling hyperparameters."""
|
||||
|
||||
top_k: int = 50
|
||||
top_p: float = 1.0
|
||||
temperature: float = 1.0
|
||||
max_tokens: int = 1024
|
||||
|
||||
|
||||
class GenerationRequest:
|
||||
"""Request parameters for text generation.
|
||||
|
||||
Encapsulates messages, sampling parameters, and streaming preference
|
||||
for a single generation request.
|
||||
Encapsulates messages, sampling parameters (via GenerationParams),
|
||||
and streaming preference for a single generation request.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -39,13 +58,31 @@ class GenerationRequest:
|
||||
stream: Whether to return output as a token stream.
|
||||
"""
|
||||
self.messages = messages
|
||||
self.top_k = top_k
|
||||
self.top_p = top_p
|
||||
self.temperature = temperature
|
||||
self.max_len = max_len
|
||||
self.params = GenerationParams(
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
temperature=temperature,
|
||||
max_tokens=max_len,
|
||||
)
|
||||
self.stream = stream
|
||||
self._validate()
|
||||
|
||||
@property
|
||||
def top_k(self) -> int:
|
||||
return self.params.top_k
|
||||
|
||||
@property
|
||||
def top_p(self) -> float:
|
||||
return self.params.top_p
|
||||
|
||||
@property
|
||||
def temperature(self) -> float:
|
||||
return self.params.temperature
|
||||
|
||||
@property
|
||||
def max_len(self) -> int:
|
||||
return self.params.max_tokens
|
||||
|
||||
def _validate(self):
|
||||
"""Validates sampling parameter ranges."""
|
||||
if not (isinstance(self.top_k, int) and self.top_k >= 0):
|
||||
@@ -296,10 +333,10 @@ class InferenceEngine:
|
||||
return self.generate(
|
||||
prompt=prompt,
|
||||
stream=request.stream,
|
||||
max_tokens=request.max_len,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
max_tokens=request.params.max_tokens,
|
||||
temperature=request.params.temperature,
|
||||
top_p=request.params.top_p,
|
||||
top_k=request.params.top_k,
|
||||
)
|
||||
|
||||
def _generate_streaming(
|
||||
|
||||
Reference in New Issue
Block a user