Spaces:
Sleeping
Sleeping
File size: 7,624 Bytes
4e316d6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 | """
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,
))
|