Anuragggggggg commited on
Commit
454c936
Β·
verified Β·
1 Parent(s): ae8ce86

Upload app.py

Browse files
Files changed (1) hide show
  1. app.py +43 -12
app.py CHANGED
@@ -76,7 +76,11 @@ if os.getenv("LAYA_HIDE_CUDA", "0") in ("1", "true", "True"):
76
  SERVING_ELSEWHERE = os.environ.get("LAYA_SERVING") == "1"
77
  BACKEND = os.getenv("LAYA_BACKEND", "").strip()
78
  MODEL = os.getenv("LAYA_MODEL", "").strip()
79
- DTYPE = os.getenv("LAYA_DTYPE", "float16").strip()
 
 
 
 
80
  TOKEN = os.getenv("LAYA_TOKEN", "").strip()
81
  ORIGINS = [o.strip() for o in os.getenv("LAYA_CORS_ORIGINS", "").split(",") if o.strip()]
82
  EAGER = os.getenv("LAYA_EAGER", "1") not in ("0", "false", "False") and not SERVING_ELSEWHERE
@@ -111,13 +115,15 @@ def agent():
111
  name = _pick_backend()
112
  mod = importlib.import_module(name)
113
  model = MODEL or DEFAULT_MODELS[name]
 
 
114
  # Both packages expose load(model_id, dtype=...); backends that do not
115
  # take dtype raise TypeError, so fall back to the plain call.
116
  try:
117
- a = mod.load(model, dtype=DTYPE)
118
  except TypeError:
119
  a = mod.load(model)
120
- _state.update(agent=a, backend=name, model=model, load_ms=round((time.time() - t0) * 1000))
121
  return a
122
  except Exception as e: # noqa: BLE001
123
  _state["error"] = f"{type(e).__name__}: {e}"
@@ -144,7 +150,11 @@ def _auth(token: Optional[str]) -> None:
144
  class Question(BaseModel):
145
  type: Literal["choice", "score", "noul"]
146
  instructions: str
147
- criteria: Optional[List[str]] = None # required for choice and score
 
 
 
 
148
 
149
 
150
  class DecideIn(BaseModel):
@@ -182,20 +192,35 @@ def _decide(state: Any, questions: Dict[str, Dict[str, Any]]) -> Dict[str, Any]:
182
 
183
 
184
  def _route(question: str, candidates: List[Candidate], instructions: str, top_k: int) -> Dict[str, Any]:
185
- labels = [c.label if not c.hint else f"{c.label} β€” {c.hint}" for c in candidates]
186
- out = _decide(question, {"pick": {"type": "choice", "instructions": instructions, "criteria": labels}})
 
 
 
 
 
 
 
 
 
 
 
 
 
187
  ans = out["answers"].get("pick", {}) if isinstance(out["answers"], dict) else {}
188
  probs = ans.get("probabilities") or ans.get("probs") or {}
189
  if isinstance(probs, list): # some builds return a bare vector
190
- probs = {labels[i]: p for i, p in enumerate(probs) if i < len(labels)}
191
  ranked = sorted(
192
- ({"id": c.id, "label": c.label, "p": float(probs.get(labels[i], 0.0))} for i, c in enumerate(candidates)),
193
  key=lambda r: r["p"], reverse=True,
194
  )[: max(1, top_k)]
195
- chosen = ans.get("value") or ans.get("label")
196
- top = next((r for r in ranked if str(r["label"]) == str(chosen)), ranked[0] if ranked else None)
 
197
  return {"id": top["id"] if top else None, "label": top["label"] if top else None,
198
- "p": top["p"] if top else None, "ranked": ranked, "ms": out["ms"]}
 
199
 
200
 
201
  # ── REST routes (attached to whichever server ends up running) ──────────────
@@ -208,6 +233,7 @@ def health() -> Dict[str, Any]:
208
  "ok": _state["agent"] is not None,
209
  "backend": _state["backend"],
210
  "model": _state["model"] or MODEL or "(default)",
 
211
  "load_ms": _state["load_ms"],
212
  "error": _state["error"],
213
  "auth_required": bool(TOKEN),
@@ -217,7 +243,12 @@ def health() -> Dict[str, Any]:
217
  @api.post("/v1/decide")
218
  def decide(body: DecideIn, x_laya_token: Optional[str] = Header(None)) -> Dict[str, Any]:
219
  _auth(x_laya_token)
220
- qs = {k: {kk: vv for kk, vv in v.model_dump().items() if vv is not None} for k, v in body.questions.items()}
 
 
 
 
 
221
  return _decide(body.state, qs)
222
 
223
 
 
76
  SERVING_ELSEWHERE = os.environ.get("LAYA_SERVING") == "1"
77
  BACKEND = os.getenv("LAYA_BACKEND", "").strip()
78
  MODEL = os.getenv("LAYA_MODEL", "").strip()
79
+ # Empty by default: the right dtype depends on the backend β€” see agent().
80
+ # float16 on a CPU produced NaNs here, which surface as a UNIFORM probability
81
+ # distribution and confidence 0.0 rather than as an error (seen 22 Sep 2026 on
82
+ # the Space: every choice 0.3333).
83
+ DTYPE = os.getenv("LAYA_DTYPE", "").strip()
84
  TOKEN = os.getenv("LAYA_TOKEN", "").strip()
85
  ORIGINS = [o.strip() for o in os.getenv("LAYA_CORS_ORIGINS", "").split(",") if o.strip()]
86
  EAGER = os.getenv("LAYA_EAGER", "1") not in ("0", "false", "False") and not SERVING_ELSEWHERE
 
115
  name = _pick_backend()
116
  mod = importlib.import_module(name)
117
  model = MODEL or DEFAULT_MODELS[name]
118
+ # MLX runs fp16 on the GPU happily; torch on a CPU needs fp32.
119
+ dtype = DTYPE or ("float16" if name == "laya_mlx" else "float32")
120
  # Both packages expose load(model_id, dtype=...); backends that do not
121
  # take dtype raise TypeError, so fall back to the plain call.
122
  try:
123
+ a = mod.load(model, dtype=dtype)
124
  except TypeError:
125
  a = mod.load(model)
126
+ _state.update(agent=a, backend=name, model=model, dtype=dtype, load_ms=round((time.time() - t0) * 1000))
127
  return a
128
  except Exception as e: # noqa: BLE001
129
  _state["error"] = f"{type(e).__name__}: {e}"
 
150
  class Question(BaseModel):
151
  type: Literal["choice", "score", "noul"]
152
  instructions: str
153
+ # Laya scores each option at its own [MASK] token, so `criteria` is a MAP
154
+ # of option -> what that option means. A bare list is accepted and each
155
+ # entry becomes its own description; passing a list to the model itself
156
+ # yields a uniform distribution with confidence 0.0 rather than an error.
157
+ criteria: Optional[Union[Dict[str, str], List[str]]] = None
158
 
159
 
160
  class DecideIn(BaseModel):
 
192
 
193
 
194
  def _route(question: str, candidates: List[Candidate], instructions: str, top_k: int) -> Dict[str, Any]:
195
+ """One choice question whose options are the candidates.
196
+
197
+ `criteria` maps each option to what it means β€” that description is what
198
+ Laya scores against the state, so a candidate's `hint` is worth giving.
199
+ Duplicate labels would collide as dict keys, so they are suffixed.
200
+ """
201
+ options: Dict[str, Candidate] = {}
202
+ for c in candidates:
203
+ key = c.label
204
+ if key in options:
205
+ key = f"{c.label} ({c.id})"
206
+ options[key] = c
207
+ criteria = {k: (c.hint or c.label) for k, c in options.items()}
208
+
209
+ out = _decide(question, {"pick": {"type": "choice", "instructions": instructions, "criteria": criteria}})
210
  ans = out["answers"].get("pick", {}) if isinstance(out["answers"], dict) else {}
211
  probs = ans.get("probabilities") or ans.get("probs") or {}
212
  if isinstance(probs, list): # some builds return a bare vector
213
+ probs = {k: probs[i] for i, k in enumerate(options) if i < len(probs)}
214
  ranked = sorted(
215
+ ({"id": c.id, "label": c.label, "p": float(probs.get(k, 0.0))} for k, c in options.items()),
216
  key=lambda r: r["p"], reverse=True,
217
  )[: max(1, top_k)]
218
+ chosen = ans.get("choice") or ans.get("value") or ans.get("label")
219
+ picked = options.get(str(chosen))
220
+ top = (next((r for r in ranked if r["id"] == picked.id), None) if picked else None) or (ranked[0] if ranked else None)
221
  return {"id": top["id"] if top else None, "label": top["label"] if top else None,
222
+ "p": top["p"] if top else None, "confidence": ans.get("confidence"),
223
+ "ranked": ranked, "ms": out["ms"]}
224
 
225
 
226
  # ── REST routes (attached to whichever server ends up running) ──────────────
 
233
  "ok": _state["agent"] is not None,
234
  "backend": _state["backend"],
235
  "model": _state["model"] or MODEL or "(default)",
236
+ "dtype": _state.get("dtype"),
237
  "load_ms": _state["load_ms"],
238
  "error": _state["error"],
239
  "auth_required": bool(TOKEN),
 
243
  @api.post("/v1/decide")
244
  def decide(body: DecideIn, x_laya_token: Optional[str] = Header(None)) -> Dict[str, Any]:
245
  _auth(x_laya_token)
246
+ qs = {}
247
+ for k, v in body.questions.items():
248
+ q = {kk: vv for kk, vv in v.model_dump().items() if vv is not None}
249
+ if isinstance(q.get("criteria"), list):
250
+ q["criteria"] = {c: c for c in q["criteria"]}
251
+ qs[k] = q
252
  return _decide(body.state, qs)
253
 
254