"""Server-side parameter validation, shared by the UI and the public /trace API.""" from __future__ import annotations DEFAULT_REVERSE = "You are a careless assistant. Give vague, unhelpful and partly incorrect answers." COMMON = { "prompt": ("text", 4000, ""), "system": ("text", 1000, ""), "raw": ("bool", False), "mode": ("choice", ("greedy", "sampling"), "greedy"), "temperature": ("float", 0.05, 2.0, 1.0), "seed": ("int", 0, 2**31 - 1, 0), "rep": ("float", 1.0, 2.0, 1.0), } SPECS = { "contrastive": { "method": ("choice", ("cd", "cad", "rose", "dola"), "cd"), "max_new_tokens": ("int", 1, 128, 64), "context": ("text", 4000, ""), "reverse_system": ("text", 1000, DEFAULT_REVERSE), "alpha": ("float", 0.0, 5.0, 0.5), "beta": ("float", 0.0, 0.9, 0.1), "tau": ("float", 0.1, 3.0, 1.0), "bucket": ("choice", ("high", "low"), "high"), "apply_norm": ("bool", True), "show_third": ("bool", True), }, "guided": { "method": ("choice", ("classifier", "regex"), "classifier"), "max_new_tokens": ("int", 1, 96, 48), "attribute": ("choice", ("sentiment", "topic"), "sentiment"), "target": ("choice", ("positive", "negative"), "positive"), "topic": ("text", 200, "astronomy"), "lam": ("float", 0.0, 50.0, 3.0), "k": ("int", 2, 40, 10), "lookahead": ("int", 0, 6, 0), "pattern": ("text", 300, r"\d{4}-\d{2}-\d{2}"), "top_k": ("int", 4, 256, 64), }, "parallel": { "method": ("choice", ("speculative", "pld", "jacobi"), "speculative"), "max_new_tokens": ("int", 1, 256, 128), "gamma": ("int", 1, 10, 4), "num_pred": ("int", 1, 20, 10), "ngram_max": ("int", 1, 5, 3), "ngram_min": ("int", 1, 5, 1), "block": ("int", 2, 16, 8), "init": ("choice", ("repeat", "random"), "repeat"), }, } def _coerce(rule: tuple, value): kind = rule[0] if kind == "text": return (rule[2] if value is None else str(value))[: rule[1]] if kind == "bool": return rule[1] if value is None else bool(value) if kind == "choice": return value if value in rule[1] else rule[2] lo, hi, default = rule[1:] try: x = float(value) except (TypeError, ValueError): x = float(default) if x != x: # NaN x = float(default) x = min(max(x, lo), hi) return int(round(x)) if kind == "int" else x def clamp_params(paradigm: str, raw: dict) -> dict: """Fill defaults, clamp ranges, and reject inputs a method can't run with.""" if paradigm not in SPECS: raise ValueError(f"Unknown paradigm {paradigm!r}.") spec = {**COMMON, **SPECS[paradigm]} out = {key: _coerce(rule, (raw or {}).get(key)) for key, rule in spec.items()} if not out["prompt"].strip(): raise ValueError("Please enter a prompt.") if paradigm == "contrastive": if out["method"] == "cad" and not out["context"].strip(): raise ValueError("Context-aware decoding needs a context passage.") if out["method"] == "rose" and not out["reverse_system"].strip(): raise ValueError("The ROSE-style contrast needs a reverse system prompt.") if paradigm == "guided" and out["method"] == "classifier" and out["attribute"] == "topic" and not out["topic"].strip(): raise ValueError("Please enter a topic for topic guidance.") if paradigm == "guided" and out["method"] == "regex" and not out["pattern"]: raise ValueError("Please enter a regular expression.") if paradigm == "parallel": out["ngram_min"] = min(out["ngram_min"], out["ngram_max"]) if out["method"] in ("pld", "jacobi"): out["mode"] = "greedy" # these methods verify greedily return out