financial-rag / src /serving /engine.py
tolivert's picture
deploy: financial_rag streamlit app
4e316d6
Raw
History Blame Contribute Delete
7.62 kB
"""
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,
))