wang2226's picture
Beyond Tokens decoding playground: contrastive, guided and parallel decoding
371d90c verified
Raw History Blame Contribute Delete
3.85 kB
"""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