Spaces:
Sleeping
Sleeping
| """ | |
| 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="<eos>", | |
| 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="<eos>", | |
| 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 <think>...</think> 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 "</think>" in buffer: | |
| after = buffer.split("</think>", 1)[1].lstrip() | |
| buffer = "" | |
| inside_think = False | |
| if after: | |
| yield after | |
| else: | |
| if "<think>" in buffer: | |
| before = buffer.split("<think>", 1)[0] | |
| inside_think = True | |
| buffer = "" | |
| if before: | |
| yield before | |
| elif _could_be_tag_prefix(buffer, "<think>"): | |
| pass # hold buffer — may still become <think> | |
| 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, | |
| )) | |