refactor: 优化工厂模式结构

This commit is contained in:
2026-04-04 11:33:58 +08:00
parent 7861af12e4
commit e97536758f
9 changed files with 164 additions and 136 deletions
+7 -6
View File
@@ -1,10 +1,11 @@
from dataclasses import dataclass
from typing import Generator, List, Optional, Tuple, Union
from typing import Dict, Generator, List, Optional, Tuple, Union
import torch
from torch import Tensor
from astrai.config.param_config import ModelParameter
from astrai.core.factory import BaseFactory
from astrai.inference.core import EmbeddingEncoderCore, GeneratorCore, KVCacheManager
HistoryType = List[Tuple[str, str]]
@@ -254,7 +255,7 @@ class EmbeddingEncoder(EmbeddingEncoderCore):
return super().encode(sentence)
class GeneratorFactory:
class GeneratorFactory(BaseFactory[GeneratorCore]):
"""Factory class for creating generator instances.
Provides smart generator selection based on request characteristics:
@@ -263,14 +264,14 @@ class GeneratorFactory:
- Single: Use LoopGenerator for single query non-streaming
Example usage:
generator = GeneratorFactory.create_generator(parameter, request)
generator = GeneratorFactory.create(parameter, request)
result = generator.generate(request)
"""
_registry: Dict[str, type] = {}
@staticmethod
def create_generator(
parameter: ModelParameter, request: GenerationRequest
) -> GeneratorCore:
def create(parameter: ModelParameter, request: GenerationRequest) -> GeneratorCore:
"""Create a generator based on request characteristics.
Args: