./llama-cli \ -m ./models/gemma-4-9b-it-Q4_K_M.gguf \ --grammar-file ./classify.gbnf \ -p "Classify this message: I want a refund for order #1234" \ -n 64 --temp 0.1
"""JSON Schema の書き方ごとに, 制約デコードの索引がどれだけ重くなるかを測る.依存なし (Python 3.10+ 標準ライブラリのみ).スキーマ相当の正規表現を組み, NFA -> DFA と落として状態数と構築時間を測る."""import timeimport statisticsclass R: """正規表現ノード. op は lit / seq / alt / opt / star.""" __slots__ = ("op", "kids", "sym") def __init__(self, op, kids=(), sym=None): self.op, self.kids, self.sym = op, tuple(kids), symdef lit(sym): return R("lit", (), sym)def opt(k): return R("opt", (k,))def star(k): return R("star", (k,))def text(s): return seq(*[lit(c) for c in s])def seq(*ks): ks = [k for k in ks if k is not None] return ks[0] if len(ks) == 1 else R("seq", ks)def alt(*ks): return ks[0] if len(ks) == 1 else R("alt", ks)# アルファベットは実文字 + 「その他の印字可能文字」1 個で近似するSTRUCT, DIGITS = set('{}[]:,"\\ '), set("0123456789")OTHER = "\x00"ALPHABET = STRUCT | DIGITS | set("abcdefghijklmnopqrstuvwxyz_-") | {OTHER}def free_string(): """任意の JSON 文字列. エスケープも受理する.""" body = star(alt(alt(*[lit(c) for c in sorted(ALPHABET - set('"\\'))]), seq(lit('\\'), alt(lit('"'), lit('\\'), lit('n'), lit('t'))))) return seq(lit('"'), body, lit('"'))def integer(): d = alt(*[lit(c) for c in sorted(DIGITS)]) return seq(opt(lit('-')), d, star(d))def enum_of(words): return alt(*[text('"%s"' % w) for w in words])def obj(pairs, additional=False): """キーはスキーマ順に出る前提 (Outlines の既定挙動に合わせる).""" inner = [] for i, (k, v) in enumerate(pairs): if i: inner.append(lit(',')) inner.append(seq(text('"%s"' % k), lit(':'), v)) body = seq(*inner) if additional: # additionalProperties: true 相当 body = seq(body, star(seq(lit(','), free_string(), lit(':'), alt(free_string(), integer())))) return seq(lit('{'), body, lit('}'))def array(item, lo, hi): """minItems=lo, maxItems=hi の配列は内部で展開される.""" parts = [] for i in range(lo): if i: parts.append(lit(',')) parts.append(item) for i in range(hi - lo): parts.append(opt(seq(lit(','), item) if (lo + i) > 0 else item)) return seq(lit('['), seq(*parts) if parts else None, lit(']'))class NFA: def __init__(self): self.t, self.e = [], [] # t: 文字遷移, e: ε 遷移 def new(self): self.t.append({}) self.e.append([]) return len(self.t) - 1def build(node, n, s): """node を NFA に展開し, 受理側の状態番号を返す.""" if node.op == "lit": d = n.new() n.t[s].setdefault(node.sym, set()).add(d) return d if node.op == "seq": for k in node.kids: s = build(k, n, s) return s if node.op == "alt": out = n.new() for k in node.kids: b = n.new() n.e[s].append(b) n.e[build(k, n, b)].append(out) return out if node.op == "opt": out = n.new() n.e[s].append(out) b = n.new() n.e[s].append(b) n.e[build(node.kids[0], n, b)].append(out) return out if node.op == "star": loop = n.new() n.e[s].append(loop) b = n.new() n.e[loop].append(b) n.e[build(node.kids[0], n, b)].append(loop) return loop raise ValueError(node.op)def closure(n, states): """ε 閉包.""" stack, seen = list(states), set(states) while stack: for d in n.e[stack.pop()]: if d not in seen: seen.add(d) stack.append(d) return frozenset(seen)def dfa_states(node): """部分集合構成で DFA の状態数を数える.""" n = NFA() start = n.new() build(node, n, start) i0 = closure(n, {start}) seen, work = {i0}, [i0] while work: moves = {} for s in work.pop(): for sym, ds in n.t[s].items(): moves.setdefault(sym, set()).update(ds) for ds in moves.values(): nxt = closure(n, ds) if nxt not in seen: seen.add(nxt) work.append(nxt) return len(seen)def measure(node, reps=5): ts = [] for _ in range(reps): t0 = time.perf_counter() states = dfa_states(node) ts.append((time.perf_counter() - t0) * 1000) return states, statistics.median(ts)LABELS = ["billing", "support", "general", "refund", "shipping", "account", "legal", "abuse", "spam", "other", "sales", "press", "bug", "feature", "docs", "trial", "renew", "cancel", "upsell", "churn"]def nested(depth, leaf): node = leaf for _ in range(depth - 1): node = obj([("child", node)]) return nodeif __name__ == "__main__": print("[1] フィールド 1 本のコスト (int vs 自由文字列)") for n in (1, 2, 3, 4, 5): a = measure(obj([("f%d" % i, integer()) for i in range(n)])) b = measure(obj([("f%d" % i, free_string()) for i in range(n)])) print(f" fields={n} int: {a[0]:4d} states {a[1]:7.2f}ms | " f"str: {b[0]:4d} states {b[1]:8.2f}ms") print("[2] 入れ子の深さ (enum + int のみ, 自由文字列なし)") leaf = obj([("label", enum_of(LABELS[:3])), ("score", integer())]) for d in (1, 3, 5, 8): s, ms = measure(nested(d, leaf)) print(f" depth={d} {s:4d} states {ms:7.2f}ms") print("[3] 開いた口・配列の上限・enum の幅") fields = [("label", enum_of(LABELS[:3])), ("note", free_string()), ("score", integer())] for name, node in (("additionalProperties:false", obj(fields)), ("additionalProperties:true ", obj(fields, additional=True))): s, ms = measure(node) print(f" {name} {s:4d} states {ms:8.2f}ms") for lo, hi in ((1, 1), (1, 3), (1, 10), (1, 20)): s, ms = measure(array(enum_of(LABELS[:3]), lo, hi)) print(f" array items={lo}..{hi:<2d} {s:4d} states {ms:8.2f}ms") for e in (3, 5, 10, 20): s, ms = measure(obj([("label", enum_of(LABELS[:e])), ("note", free_string())])) print(f" enum={e:<2d} choices {s:4d} states {ms:8.2f}ms")
"""分類器の CI テスト. 温めたモデルに対して数秒で終わる."""import pytestfrom agents.classifier import classify@pytest.mark.parametrize("message,expected", [ ("I want a refund for order #1234", "billing"), ("Where is my order?", "support"), ("Hello, how are you?", "general"),])def test_classify(message: str, expected: str) -> None: assert classify(message) == expected