File size: 3,431 Bytes
4fdb3ad
93b1c3a
 
 
4fdb3ad
 
 
 
 
 
 
 
93b1c3a
 
 
 
 
4fdb3ad
 
 
93b1c3a
4fdb3ad
 
93b1c3a
4fdb3ad
 
 
 
3c3e911
 
 
 
 
 
 
 
 
4fdb3ad
93b1c3a
 
 
 
 
 
ac4763f
 
 
 
93b1c3a
 
4fdb3ad
3c3e911
4fdb3ad
 
 
 
 
3c3e911
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
model.py — Text generation.
Uses Modal GPU (Qwen3-4B) when MODAL_TOKEN_ID/SECRET are set,
falls back to local llama-cpp GGUF otherwise.
"""
import os
from functools import lru_cache

MODEL_REPO = os.getenv("MODEL_REPO", "Qwen/Qwen3-1.7B-GGUF")
MODEL_FILE = os.getenv("MODEL_FILE", "Qwen3-1.7B-Q8_0.gguf")
N_CTX = int(os.getenv("N_CTX", "4096"))
N_THREADS = int(os.getenv("N_THREADS", str(os.cpu_count() or 4)))
MODAL_APP = os.getenv("MODAL_APP_NAME", "storyforge")


def _modal_ready() -> bool:
    return bool(os.getenv("MODAL_TOKEN_ID") and os.getenv("MODAL_TOKEN_SECRET"))


@lru_cache(maxsize=1)
def _local_llm():
    from huggingface_hub import hf_hub_download
    from llama_cpp import Llama

    path = hf_hub_download(repo_id=MODEL_REPO, filename=MODEL_FILE)
    return Llama(model_path=path, n_ctx=N_CTX, n_threads=N_THREADS, verbose=False)


def _local_messages(system: str, user: str) -> list:
    # "/no_think" is Qwen3's soft switch — without it the local GGUF burns most
    # of the token budget inside <think> blocks and truncates the JSON.
    return [
        {"role": "system", "content": system + " /no_think"},
        {"role": "user", "content": user},
    ]


def generate(system: str, user: str, max_tokens: int = 512) -> str:
    if _modal_ready():
        try:
            import modal

            TextModel = modal.Cls.from_name(MODAL_APP, "TextModel")
            return TextModel().generate.remote(system, user, max_tokens)
        except Exception as e:
            import traceback
            print(f"[model] Modal call failed, falling back to local: {e}")
            traceback.print_exc()
    # Local fallback
    llm = _local_llm()
    out = llm.create_chat_completion(
        messages=_local_messages(system, user),
        max_tokens=max_tokens,
        temperature=0.8,
        top_p=0.9,
    )
    return out["choices"][0]["message"]["content"]


def generate_stream(system: str, user: str, max_tokens: int = 512):
    """Yield the accumulated response text as it is generated.

    Fallback chain: Modal streaming → Modal non-streaming (older deployment
    without generate_stream) → local llama-cpp streaming.
    """
    if _modal_ready():
        acc = ""
        try:
            import modal

            TextModel = modal.Cls.from_name(MODAL_APP, "TextModel")
            for piece in TextModel().generate_stream.remote_gen(system, user, max_tokens):
                acc += piece
                yield acc
            return
        except Exception as e:
            print(f"[model] Modal stream failed: {e}")
            if acc:
                # Partial stream already surfaced — restart cleanly below.
                acc = ""
        try:
            import modal

            TextModel = modal.Cls.from_name(MODAL_APP, "TextModel")
            yield TextModel().generate.remote(system, user, max_tokens)
            return
        except Exception as e:
            import traceback
            print(f"[model] Modal call failed, falling back to local: {e}")
            traceback.print_exc()
    llm = _local_llm()
    acc = ""
    for part in llm.create_chat_completion(
        messages=_local_messages(system, user),
        max_tokens=max_tokens,
        temperature=0.8,
        top_p=0.9,
        stream=True,
    ):
        delta = part["choices"][0].get("delta", {}).get("content") or ""
        if delta:
            acc += delta
            yield acc