Bc-AI commited on
Commit
3bb28fd
Β·
verified Β·
1 Parent(s): 3211cdc

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +161 -0
app.py ADDED
@@ -0,0 +1,161 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gradio as gr
2
+ from gradio import Server
3
+ import spaces
4
+ import torch
5
+ from transformers import AutoTokenizer, AutoModelForCausalLM
6
+
7
+ # ─────────────────────────────────────────────
8
+ # 1. MODEL SETUP
9
+ # ─────────────────────────────────────────────
10
+ MODEL_ID = "Smilyai-labs/Mira-1-large"
11
+
12
+ print(f"Loading tokenizer: {MODEL_ID}")
13
+ tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True)
14
+
15
+ print(f"Loading model: {MODEL_ID}")
16
+ model = AutoModelForCausalLM.from_pretrained(
17
+ MODEL_ID,
18
+ torch_dtype=torch.float16,
19
+ device_map="auto", # ZeroGPU manages CUDA device
20
+ trust_remote_code=True, # Needed for Qwen-based custom archs
21
+ )
22
+ model.eval()
23
+
24
+ # ─────────────────────────────────────────────
25
+ # 2. INFERENCE FUNCTION (ZeroGPU decorated)
26
+ # ─────────────────────────────────────────────
27
+ @spaces.GPU(duration=120)
28
+ def generate(
29
+ prompt: str,
30
+ system_prompt: str = "You are a helpful assistant.",
31
+ max_new_tokens: int = 512,
32
+ temperature: float = 0.7,
33
+ top_p: float = 0.9,
34
+ do_sample: bool = True,
35
+ ) -> str:
36
+ """
37
+ Generate a text response from Mira-1-Large.
38
+
39
+ Args:
40
+ prompt: The user message / prompt to send to the model.
41
+ system_prompt: System-level instruction for the model.
42
+ max_new_tokens: Maximum number of tokens to generate.
43
+ temperature: Sampling temperature (higher = more creative).
44
+ top_p: Nucleus sampling probability mass.
45
+ do_sample: Whether to use sampling (True) or greedy decoding (False).
46
+
47
+ Returns:
48
+ The model's text response as a string.
49
+ """
50
+ # Build chat-style messages (Qwen uses apply_chat_template)
51
+ messages = [
52
+ {"role": "system", "content": system_prompt},
53
+ {"role": "user", "content": prompt},
54
+ ]
55
+
56
+ # Qwen / Mira chat template
57
+ text = tokenizer.apply_chat_template(
58
+ messages,
59
+ tokenize=False,
60
+ add_generation_prompt=True,
61
+ )
62
+
63
+ inputs = tokenizer(text, return_tensors="pt").to(model.device)
64
+
65
+ with torch.no_grad():
66
+ output_ids = model.generate(
67
+ **inputs,
68
+ max_new_tokens=max_new_tokens,
69
+ temperature=temperature,
70
+ top_p=top_p,
71
+ do_sample=do_sample,
72
+ pad_token_id=tokenizer.eos_token_id,
73
+ )
74
+
75
+ # Decode only the newly generated tokens
76
+ new_tokens = output_ids[0][inputs["input_ids"].shape[1]:]
77
+ response = tokenizer.decode(new_tokens, skip_special_tokens=True)
78
+ return response
79
+
80
+
81
+ # ─────────────────────────────────────────────
82
+ # 3. STREAMING INFERENCE (SSE / token-by-token)
83
+ # ─────────────────────────────────────────────
84
+ @spaces.GPU(duration=120)
85
+ def generate_stream(
86
+ prompt: str,
87
+ system_prompt: str = "You are a helpful assistant.",
88
+ max_new_tokens: int = 512,
89
+ temperature: float = 0.7,
90
+ top_p: float = 0.9,
91
+ ) -> str:
92
+ """
93
+ Stream a text response token-by-token from Mira-1-Large via SSE.
94
+
95
+ Args:
96
+ prompt: The user message / prompt.
97
+ system_prompt: System-level instruction for the model.
98
+ max_new_tokens: Maximum number of tokens to generate.
99
+ temperature: Sampling temperature.
100
+ top_p: Nucleus sampling probability mass.
101
+
102
+ Yields:
103
+ Partial response strings, growing with each new token.
104
+ """
105
+ from transformers import TextIteratorStreamer
106
+ from threading import Thread
107
+
108
+ messages = [
109
+ {"role": "system", "content": system_prompt},
110
+ {"role": "user", "content": prompt},
111
+ ]
112
+ text = tokenizer.apply_chat_template(
113
+ messages, tokenize=False, add_generation_prompt=True
114
+ )
115
+ inputs = tokenizer(text, return_tensors="pt").to(model.device)
116
+
117
+ streamer = TextIteratorStreamer(
118
+ tokenizer, skip_prompt=True, skip_special_tokens=True
119
+ )
120
+
121
+ gen_kwargs = dict(
122
+ **inputs,
123
+ max_new_tokens=max_new_tokens,
124
+ temperature=temperature,
125
+ top_p=top_p,
126
+ do_sample=True,
127
+ streamer=streamer,
128
+ pad_token_id=tokenizer.eos_token_id,
129
+ )
130
+
131
+ thread = Thread(target=model.generate, kwargs=gen_kwargs)
132
+ thread.start()
133
+
134
+ partial = ""
135
+ for new_text in streamer:
136
+ partial += new_text
137
+ yield partial
138
+
139
+
140
+ # ─────────────────────────────────────────────
141
+ # 4. gr.Server β€” REST API + OPTIONAL SWAGGER UI
142
+ # ───────────────────��─────────────────────────
143
+ app = Server(
144
+ title="Mira-1-Large API",
145
+ summary="ZeroGPU-backed REST API for Smilyai-labs/Mira-1-large (Qwen arch)",
146
+ version="1.0.0",
147
+ )
148
+
149
+ # Register as Gradio API endpoints (queued, SSE-streaming capable)
150
+ app.api(generate, name="generate") # POST /gradio_api/call/generate
151
+ app.api(generate_stream, name="generate_stream") # POST /gradio_api/call/generate_stream
152
+
153
+ # Optional: plain FastAPI GET health-check route
154
+ @app.get("/health")
155
+ def health():
156
+ return {"status": "ok", "model": MODEL_ID}
157
+
158
+ # ─────────────────────────────────────────────
159
+ # 5. LAUNCH
160
+ # ─────────────────────────────────────────────
161
+ app.launch()