""" Generation engine wrapping a transformer model + tokenizer. Supports multiple architectures (Qwen3, Gemma3, Gemma4) via the ``arch`` parameter. Loads model weights once at startup and exposes a streaming generate method that yields decoded text chunks. All PyTorch inference runs inside this class so the FastAPI layer stays framework-agnostic. """ import torch from src.generation.text_generator import generate_text_basic_stream from src.utils.device import get_default_device # --------------------------------------------------------------------------- # Architecture registry — maps arch name to (config_fn, model_cls, loader_fn, # tokenizer_cls, tokenizer_repo_fn, eos_key, loader_kwargs_fn) # --------------------------------------------------------------------------- def _qwen3_loader_kwargs(model_size: str) -> dict: return {"model_size": model_size, "use_reasoning": True} def _gemma3_loader_kwargs(model_size: str) -> dict: return {"model_size": model_size, "use_reasoning": True} def _gemma4_loader_kwargs(_model_size: str) -> dict: return {} _ARCH_REGISTRY: dict[str, dict] = {} def _register_arch( name: str, config_import: tuple[str, str], # (module, function) model_import: tuple[str, str], # (module, class) loader_import: tuple[str, str], # (module, function) tokenizer_import: tuple[str, str], # (module, class) tokenizer_repo_fn, # model_size -> repo string eos_key: str, # key to look up in tokenizer._special_to_id loader_kwargs_fn=None, # model_size -> extra kwargs for loader default_size: str = "0.6B", ): _ARCH_REGISTRY[name] = { "config_import": config_import, "model_import": model_import, "loader_import": loader_import, "tokenizer_import": tokenizer_import, "tokenizer_repo_fn": tokenizer_repo_fn, "eos_key": eos_key, "loader_kwargs_fn": loader_kwargs_fn or (lambda s: {"model_size": s}), "default_size": default_size, } _register_arch( "qwen3", config_import=("src.architectures.qwen3.config", "get_config"), model_import=("src.architectures.qwen3.model", "Qwen3Model"), loader_import=("src.architectures.qwen3.loader", "download_and_load_qwen"), tokenizer_import=("src.tokenization.qwen_tokenizer", "Qwen3Tokenizer"), tokenizer_repo_fn=lambda s: f"Qwen/Qwen3-{s}", eos_key="<|im_end|>", loader_kwargs_fn=_qwen3_loader_kwargs, default_size="0.6B", ) _register_arch( "gemma3", config_import=("src.architectures.gemma3.config", "get_config"), model_import=("src.architectures.gemma3.model", "Gemma3Model"), loader_import=("src.architectures.gemma3.loader", "download_and_load_gemma3"), tokenizer_import=("src.tokenization.gemma3_tokenizer", "Gemma3Tokenizer"), tokenizer_repo_fn=lambda s: f"google/gemma-3-{s}-it", eos_key="", loader_kwargs_fn=_gemma3_loader_kwargs, default_size="270m", ) _register_arch( "gemma4", config_import=("src.architectures.gemma4.config", "get_config"), model_import=("src.architectures.gemma4.model", "Gemma4Model"), loader_import=("src.architectures.gemma4.loader", "download_and_load_gemma4"), tokenizer_import=("src.tokenization.gemma4_tokenizer", "Gemma4Tokenizer"), tokenizer_repo_fn=lambda _s: "google/gemma-4-E2B", eos_key="", loader_kwargs_fn=_gemma4_loader_kwargs, default_size="E2B", ) def _import_attr(module_path: str, attr_name: str): """Lazily import an attribute from a module.""" import importlib mod = importlib.import_module(module_path) return getattr(mod, attr_name) def _could_be_tag_prefix(text: str, tag: str) -> bool: """Check if `text` ends with a string that is a prefix of `tag`.""" for i in range(1, len(tag)): if text.endswith(tag[:i]): return True return False class GenerationEngine: """Architecture-agnostic generation engine. Parameters ---------- model_size : str Size variant within the architecture (e.g. "0.6B", "270m", "E2B"). arch : str Architecture family: ``"qwen3"`` (default), ``"gemma3"``, or ``"gemma4"``. """ def __init__(self, model_size: str | None = None, arch: str = "qwen3"): self.device = get_default_device() self.arch = arch if arch not in _ARCH_REGISTRY: raise ValueError( f"Unknown architecture: {arch}. " f"Choose from {list(_ARCH_REGISTRY.keys())}" ) reg = _ARCH_REGISTRY[arch] self.model_size = model_size or reg["default_size"] # Lazy imports — only pull in the architecture we need get_config = _import_attr(*reg["config_import"]) ModelClass = _import_attr(*reg["model_import"]) load_fn = _import_attr(*reg["loader_import"]) TokenizerClass = _import_attr(*reg["tokenizer_import"]) # Build model cfg = get_config(self.model_size) self.model = ModelClass(cfg) loader_kwargs = reg["loader_kwargs_fn"](self.model_size) load_fn(self.model, cfg, **loader_kwargs, device=self.device) self.model.eval() # Build tokenizer repo = reg["tokenizer_repo_fn"](self.model_size) self.tokenizer = TokenizerClass(repo) self.eos_token_id = self.tokenizer._special_to_id.get(reg["eos_key"]) def generate_stream( self, messages: list[dict], max_tokens: int = 512, temperature: float = 0.7, top_k: int = 20, top_p: float = 0.95, ): """Yield decoded text chunks, with ... blocks stripped.""" token_ids = self.tokenizer.encode(messages, add_generation_prompt=True) input_tensor = torch.tensor([token_ids], device=self.device) inside_think = False buffer = "" for next_token in generate_text_basic_stream( self.model, input_tensor, max_new_tokens=max_tokens, temperature=temperature, top_k=top_k, top_p=top_p, eos_token_id=self.eos_token_id, ): token_id = next_token[0, 0].item() text = self.tokenizer.decode([token_id]) buffer += text if inside_think: if "" in buffer: after = buffer.split("", 1)[1].lstrip() buffer = "" inside_think = False if after: yield after else: if "" in buffer: before = buffer.split("", 1)[0] inside_think = True buffer = "" if before: yield before elif _could_be_tag_prefix(buffer, ""): pass # hold buffer — may still become else: yield buffer buffer = "" if buffer and not inside_think: yield buffer def generate_answer( self, messages: list[dict], max_tokens: int = 512, temperature: float = 0.7, top_k: int = 20, top_p: float = 0.95, ) -> str: """Generate a complete response (non-streaming). Convenience wrapper for evaluation.""" return "".join(self.generate_stream( messages, max_tokens=max_tokens, temperature=temperature, top_k=top_k, top_p=top_p, ))