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,
        ))