File size: 11,130 Bytes
116c073
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f88b204
 
 
116c073
b416a56
116c073
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b416a56
116c073
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
68c6048
116c073
 
 
 
 
 
68c6048
116c073
 
 
 
68c6048
 
 
 
116c073
 
 
 
 
 
 
 
68c6048
116c073
 
 
 
 
 
 
 
68c6048
116c073
2c998cb
 
116c073
 
 
 
 
 
 
 
2c998cb
 
 
 
 
116c073
 
 
68c6048
116c073
2c998cb
116c073
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c2bb655
99e2f28
c2bb655
116c073
b416a56
116c073
 
 
 
 
2c998cb
116c073
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b416a56
116c073
 
 
 
 
b416a56
116c073
 
 
 
 
 
 
 
 
b416a56
 
116c073
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
import gc
import os
import tempfile
from collections.abc import Generator
from functools import wraps
from pathlib import Path
from threading import Lock, Thread

import spaces
import torch
import transformers
from fastapi.responses import HTMLResponse
from gradio import Server
from transformers import AutoModelForCausalLM, AutoTokenizer, LogitsProcessor, LogitsProcessorList, TextIteratorStreamer

print(f"Transformers version: {transformers.__version__}")

HF_TOKEN = os.environ.get("HF_TOKEN")

# Available model variants. The front-end populates its selector from the
# /models endpoint, so adding a new entry here is enough to offer it.
# Only one variant is kept loaded at a time; switching unloads the previous.
MODEL_VARIANTS = {
    "Nalu 1.0 (Base)": {"repo_id": "khtsly/Nalu-1.0-Base", "tokenizer": "khtsly/luau-coder-1.0-preview-tokenizer", "enabled": True},
    "Moana 1.0 (Base)": {"repo_id": "khtsly/Moana-1.0-Base", "tokenizer": "khtsly/luau-coder-1.0-preview-tokenizer", "enabled": True},
    "Moana 1.5": {"repo_id": "khtsly/Moana-1.5-Base", "tokenizer": "khtsly/luau-coder-1.5-tokenizer", "enabled": False},
}
DEFAULT_VARIANT = "Nalu 1.0 (Base)"


def _repair_checkpoint_keys(weights_path: Path) -> bool:
    """Remove an obsolete composite-model prefix before Transformers loads."""
    from safetensors import safe_open

    with safe_open(weights_path, framework="pt", device="cpu") as f:
        keys = list(f.keys())

    prefixed = [key for key in keys if key.startswith("language_model.")]
    if not prefixed:
        return False
    if len(prefixed) != len(keys):
        raise RuntimeError(
            f"Refusing to repair mixed checkpoint keys in {weights_path}: "
            f"{len(prefixed)} of {len(keys)} use the language_model. prefix"
        )

    fixed_keys = [key.removeprefix("language_model.") for key in keys]
    if len(set(fixed_keys)) != len(fixed_keys):
        raise RuntimeError(f"Prefix removal would create duplicate keys in {weights_path}")

    from safetensors.torch import load_file, save_file

    print(f"Repairing {len(keys)} checkpoint keys in {weights_path} ...")
    state = load_file(weights_path, device="cpu")
    fixed = {
        key.removeprefix("language_model."): tensor.contiguous()
        for key, tensor in state.items()
    }
    tmp_path = weights_path.with_name(f".{weights_path.name}.repair.tmp")
    try:
        save_file(fixed, tmp_path)
        os.replace(tmp_path, weights_path)
    finally:
        tmp_path.unlink(missing_ok=True)
    print("Checkpoint key repair complete")
    return True


def _prepare_model_source(variant: str) -> str:
    """Return a local model directory whose safetensors keys are loadable."""
    repo_id = MODEL_VARIANTS[variant]["repo_id"]
    cache_dir = Path(tempfile.gettempdir()) / f"model-{variant}"

    # Local override for the default variant (checked-in model.safetensors).
    if variant == DEFAULT_VARIANT:
        app_dir = Path(__file__).resolve().parent
        weights_path = app_dir / "model.safetensors"
        if weights_path.is_file():
            _repair_checkpoint_keys(weights_path)
            return str(app_dir)

    from huggingface_hub import snapshot_download

    snapshot_download(
        repo_id=repo_id,
        local_dir=cache_dir,
        token=HF_TOKEN,
        allow_patterns=["*.json", "*.py", "*.jinja", "*.safetensors"],
    )
    weights_path = cache_dir / "model.safetensors"

    if not weights_path.is_file():
        raise FileNotFoundError(f"No model.safetensors found in {cache_dir}")

    _repair_checkpoint_keys(weights_path)
    return str(cache_dir)

_model = None
_tokenizer = None
_current_variant = None
_model_lock = Lock()


def _unload_current_model() -> None:
    """Free the currently loaded variant so the next one can take its place."""
    global _model, _current_variant, _tokenizer
    if _model is not None:
        print(f"Unloading model variant: {_current_variant}")
        del _model
        _model = None
    if _tokenizer is not None:
        print(f"Unloading tokenizer variant: {_current_variant}")
        del _tokenizer
        _tokenizer = None
    _current_variant = None
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()


def ensure_model(variant: str):
    """Load `variant` on demand, unloading whatever is currently loaded."""
    global _model, _current_variant, _tokenizer
    if variant not in MODEL_VARIANTS:
        print(f"Unknown model variant {variant!r}; falling back to {DEFAULT_VARIANT}")
        variant = DEFAULT_VARIANT
    if not MODEL_VARIANTS[variant].get("enabled", True):
        print(f"Model variant {variant!r} is not enabled yet; falling back to {DEFAULT_VARIANT}")
        variant = DEFAULT_VARIANT
    with _model_lock:
        if _model is not None and _current_variant == variant:
            return _model, _tokenizer
        _unload_current_model()
        model_targeted = MODEL_VARIANTS[variant]
        print(f"Loading model variant: {variant} ({model_targeted["repo_id"]}) | tokenizer: {model_targeted["tokenizer"]} ...")
        source = _prepare_model_source(variant)
        loaded = AutoModelForCausalLM.from_pretrained(
            source,
            trust_remote_code=True,
            dtype=torch.bfloat16,
            token=HF_TOKEN,
            device_map="cuda",
        )
        tokenizer = AutoTokenizer.from_pretrained(
            model_targeted["tokenizer"],
            trust_remote_code=True,
            token=HF_TOKEN,
        )
        _install_transformers_516_cache_compat(loaded)
        _model = loaded
        _current_variant = variant
        _tokenizer = tokenizer
        print(f"Model variant ready: {variant}")
        return _model, tokenizer


def _install_transformers_516_cache_compat(loaded_model) -> None:
    """Handle empty generation positions with older Kimi remote code."""
    inner_model = loaded_model.model
    original_forward = inner_model.forward

    @wraps(original_forward)
    def compatible_forward(*args, **kwargs):
        for name in ("cache_position", "position_ids"):
            value = kwargs.get(name)
            numel = getattr(value, "numel", None)
            if callable(numel) and numel() == 0:
                kwargs[name] = None
        return original_forward(*args, **kwargs)

    inner_model.forward = compatible_forward


class PresencePenaltyLogitsProcessor(LogitsProcessor):
    """OpenAI-style presence penalty for HF generate (which lacks it natively).

    Subtracts `penalty` once from the logit of every unique token already
    present in the context, discouraging the model from reusing tokens.
    """

    def __init__(self, penalty: float):
        self.penalty = float(penalty)

    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:
        if self.penalty == 0:
            return scores
        for batch_idx, ids in enumerate(input_ids):
            for token_id in set(ids.tolist()):
                if 0 <= token_id < scores.shape[-1]:
                    scores[batch_idx, token_id] -= self.penalty
        return scores


demo = Server()


@demo.api()
@spaces.GPU(duration=60)
def predict(
    message: str,
    history: list[list] | None = None,
    max_tokens: int = 32,
    temperature: float = 0.6,
    top_p: float = 0.95,
    top_k: int = 40,
    min_p: float = 0.05,
    repetition_penalty: float = 1.05,
    presence_penalty: float = 0.0,
    model_id: str = "Nalu 1.0 (Base)",
) -> Generator[str, None, None]:
    # History is accepted for front-end chat compatibility; generation uses
    # the current message only to preserve base-model continuation behavior.
    raw_input_text = message
    max_tokens = int(max_tokens)
    active_model, tokenizer = ensure_model(model_id)
    if next(active_model.parameters()).device.type != "cuda":
        active_model.to("cuda")
    inputs = tokenizer(raw_input_text, return_tensors="pt")
    token_count = inputs["input_ids"].shape[1]
    print(f"Generation request: {len(raw_input_text)} chars, {token_count} tokens")

    if token_count == 0:
        seed_token_id = tokenizer.bos_token_id
        if seed_token_id is None:
            candidate = tokenizer.convert_tokens_to_ids("[BOS]")
            if isinstance(candidate, int) and candidate >= 0:
                seed_token_id = candidate
        if seed_token_id is None:
            yield "Please enter a non-empty prompt."
            return
        inputs["input_ids"] = torch.tensor([[seed_token_id]], dtype=torch.long)
        inputs["attention_mask"] = torch.ones((1, 1), dtype=torch.long)
        token_count = 1
        print(f"Empty prompt: seeded generation with token {seed_token_id}")

    inputs = inputs.to("cuda")
    streamer = TextIteratorStreamer(tokenizer, skip_prompt=True)

    generation_kwargs = dict(
        **inputs,
        streamer=streamer,
        max_new_tokens=int(max_tokens),
        pad_token_id=tokenizer.eos_token_id,
        use_cache=True,
        generation_mode=True,
    )

    temperature = float(temperature)
    top_p = float(top_p)
    top_k = int(top_k)
    min_p = float(min_p)
    repetition_penalty = float(repetition_penalty)
    presence_penalty = float(presence_penalty)

    if temperature > 0.0:
        generation_kwargs.update(do_sample=True, temperature=temperature)
        if 0.0 < top_p < 1.0:
            generation_kwargs["top_p"] = top_p
        if top_k > 0:
            generation_kwargs["top_k"] = top_k
        if 0.0 < min_p <= 1.0:
            generation_kwargs["min_p"] = min_p
    else:
        generation_kwargs.update(do_sample=False)

    if repetition_penalty != 1.0:
        generation_kwargs["repetition_penalty"] = repetition_penalty

    if presence_penalty != 0.0:
        generation_kwargs["logits_processor"] = LogitsProcessorList(
            [PresencePenaltyLogitsProcessor(presence_penalty)]
        )

    generation_error = []

    def run_generation():
        try:
            active_model.generate(**generation_kwargs)
        except Exception as exc:
            generation_error.append(exc)
            streamer.end()

    thread = Thread(target=run_generation)
    thread.start()

    output_buffer = ""
    for new_text in streamer:
        output_buffer += new_text
        yield output_buffer
    thread.join()

    if generation_error:
        exc = generation_error[0]
        message_out = f"Generation failed: {type(exc).__name__}: {exc}"
        print(message_out)
        yield f"{output_buffer}\n\n{message_out}" if output_buffer else message_out


@demo.get("/models")
async def list_models():
    return {
        "default": DEFAULT_VARIANT,
        "variants": [
            {"id": name, "enabled": bool(cfg.get("enabled", True))}
            for name, cfg in sorted(MODEL_VARIANTS.items())
        ],
    }


@demo.get("/", response_class=HTMLResponse)
async def homepage():
    html_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "index.html")
    with open(html_path, "r", encoding="utf-8") as f:
        return f.read()


if __name__ == "__main__":
    demo.launch(show_error=True)