diff --git a/tools/parallel-decision/README.md b/tools/parallel-decision/README.md index 76e86b90e773..46ad7de92e6e 100644 --- a/tools/parallel-decision/README.md +++ b/tools/parallel-decision/README.md @@ -110,6 +110,7 @@ Compact fields, or a JSON Schema object with `properties`: | `boolean` | - | true / false | | `integer` | `minimum`, `maximum` | 1-255 values | | `number` | `minimum`, `maximum`, `step` (`multipleOf` in JSON Schema) | fixed-width decimals | +| `string` | `max_tokens` | open field: free text, generated (see below) | Numeric fields take `aggregate`: `mode` (default), `median` or `mean`. @@ -121,6 +122,44 @@ Numeric fields take `aggregate`: `mode` (default), `median` or `mean`. | `mode` | `auto` | `tree` scores every divergence node and returns exact probabilities; `greedy` walks the trie; `auto` picks tree up to `tree_max` values | | `tree_max` | 128 | per-field switch between tree and greedy | | `cache_prompt` | true | reuse the cached instructions + schema prefix | +| `open_sampling` | `greedy` | open fields only: `greedy` (argmax, deterministic) or `temperature` | +| `open_temp` | 0.7 | open fields only: temperature when `open_sampling` is `temperature` | + +### Hybrid decisions (closed fields + one open field) + +A schema may mix closed fields with **at most one open field** (`{"type": "string", "max_tokens": N}`, N in 1-1024). +The closed fields are scored as usual; then the open field is **generated** on the same context, in the same call: +the scored closed values form the generation prefix, and the model continues autoregressively until end-of-generation +or `max_tokens`. One HTTP call, one prefill. + +```bash +curl http://localhost:8096/v1/decision -H "Content-Type: application/json" -d '{ + "model": "gemma-4-12b", + "schema": { + "category": {"type": "enum", "choices": ["billing","technical","other"], + "description": "What type of support request is this?"}, + "urgent": {"type": "boolean", "description": "Does this need urgent handling?"}, + "reason": {"type": "string", "max_tokens": 128, "description": "One-sentence reason."} + }, + "contexts": ["I was charged twice and need this fixed today."] +}' +``` + +```json +{ + "decision": {"category": "billing", "urgent": true, "reason": "The user was charged twice."}, + "fields": { + "category": {"value": "billing", "probability": 1.0, "scored_nodes": 1, "tree": true}, + "urgent": {"value": true, "probability": 1.0, "scored_nodes": 1, "tree": true}, + "reason": {"value": "The user was charged twice.", "generated": true, "tokens": 9, "truncated": false} + }, + "usage": {"context_tokens": 21, "scored_rows": 14, "generated_tokens": 9} +} +``` + +The open field's entry carries `generated: true`, `tokens` (how many were generated) and `truncated` (true when +`max_tokens` was reached before end-of-generation). `timings` gains `generation_ms`; the batch `usage` gains +`generated_tokens`. ## CLI diff --git a/tools/parallel-decision/decision-engine.cpp b/tools/parallel-decision/decision-engine.cpp index c0f1d65000c5..3f4012bd55c5 100644 --- a/tools/parallel-decision/decision-engine.cpp +++ b/tools/parallel-decision/decision-engine.cpp @@ -8,6 +8,7 @@ #include #include #include +#include #include namespace llama_decision { @@ -169,6 +170,73 @@ struct decision_field { } }; +// Greedy (argmax) or temperature sampling over the full vocabulary. +llama_token sample_token(const llama_vocab * vocab, const float * logits, const std::string & sampling, + float temp, std::mt19937 & rng, std::vector & scratch) { + const int n = llama_vocab_n_tokens(vocab); + if (sampling != "temperature" || temp <= 0.0f) { + int best = 0; + for (int i = 1; i < n; ++i) { + if (logits[i] > logits[best]) { + best = i; + } + } + return (llama_token) best; + } + float mx = logits[0]; + for (int i = 1; i < n; ++i) { + mx = std::max(mx, logits[i]); + } + scratch.resize((size_t) n); + for (int i = 0; i < n; ++i) { + scratch[i] = std::exp((logits[i] - mx) / temp); + } + std::discrete_distribution dist(scratch.begin(), scratch.end()); + return (llama_token) dist(rng); +} + +// The open field is the last field of the JSON answer: the model writes its value as a JSON +// string literal followed by the closing brace. Recover the value from the generated text. +// Primary: json_head ("{\n" + the scored closed fields + the open suffix) plus the generated +// tail is the full JSON object; parse it and read the field. Fallback: strip surrounding +// whitespace and quotes, then parse the string literal. +std::string recover_open_text(std::string text, bool truncated, const std::string & json_head, const std::string & field_name) { + if (!truncated && !field_name.empty()) { + try { + const auto obj = common_json::parse(json_head + text); + if (obj.is_object() && obj.contains(field_name) && obj.at(field_name).is_string()) { + return obj.at(field_name).get(); + } + } catch (...) { + // the tail is not a complete JSON object; fall through to literal recovery + } + } + size_t start = 0; + while (start < text.size() && (text[start] == ' ' || text[start] == '\t' || text[start] == '\n' || text[start] == '\r')) { + ++start; + } + text = text.substr(start); + if (!text.empty() && text[0] == '"') { + text.erase(0, 1); + if (!truncated) { + const size_t q = text.rfind('"'); + if (q != std::string::npos) { + const std::string lit = text.substr(0, q); + try { + const auto parsed = common_json::parse(std::string("\"") + lit + "\""); + if (parsed.is_string()) { + return parsed.get(); + } + } catch (...) { + // the model escaped something the parser rejects; keep the raw text + } + text = lit; + } + } + } + return text; +} + } // namespace // ---------------------------------------------------------------- engine @@ -314,6 +382,7 @@ result engine::decide(const std::string & shared_text, const std::string & conte r.rounds = b.rounds; r.prefill_ms = b.prefill_ms; r.scoring_ms = b.scoring_ms; + r.generation_ms = b.generation_ms; return r; } @@ -335,9 +404,25 @@ batch_result engine::decide_batch(const std::string & shared_text, const std::ve } std::vector fields; - int total = 0; - int branches = 0; // round-1 branches of one context - for (const auto & in : inputs) { + int total = 0; + int branches = 0; // round-1 branches of one context + int open_idx = -1; + for (size_t f = 0; f < inputs.size(); ++f) { + const auto & in = inputs[f]; + if (in.candidates.empty()) { + // open field: generated on the trunk after the closed fields are scored; keep a + // placeholder so the field indices stay aligned with the schema + if (in.max_tokens < 1) { + throw std::invalid_argument("an open field needs max_tokens"); + } + if (open_idx >= 0) { + throw std::invalid_argument("at most one open (free-text) field is allowed"); + } + open_idx = (int) f; + fields.emplace_back(tokens_t{}, std::vector{}); + fields.back().use_tree = false; + continue; + } const size_t n = in.candidates.size(); if (n < 1 || n > 255) { throw std::invalid_argument("each field needs 1-255 allowed values"); @@ -484,8 +569,8 @@ batch_result engine::decide_batch(const std::string & shared_text, const std::ve } first = false; } + out.scoring_ms += ms_since(ts); for (size_t i = 0; i < n_group; ++i) { - llama_memory_seq_rm(mem, seq_pool + (llama_seq_id) i, -1, -1); result & r = out.items[g0 + i]; r.context_tokens = prefixes[g0 + i].size(); r.rows = total; @@ -496,11 +581,75 @@ batch_result engine::decide_batch(const std::string & shared_text, const std::ve r.fields.push_back({ fd.winner, fd.path_score, fd.scored_nodes, fd.use_tree, fd.probs }); } } - out.scoring_ms += ms_since(ts); + // the open field is generated on each trunk before it is released: the scored closed + // fields form the generation prefix, the model continues until EOS or max_tokens + if (open_idx >= 0) { + const auto tg = std::chrono::steady_clock::now(); + for (size_t i = 0; i < n_group; ++i) { + generate_open(seq_pool + (llama_seq_id) i, (llama_pos) (shared.size() + prefixes[g0 + i].size()), + inputs, open_idx, out.items[g0 + i].fields, opt, out.items[g0 + i]); + } + out.generation_ms += ms_since(tg); + } + for (size_t i = 0; i < n_group; ++i) { + llama_memory_seq_rm(mem, seq_pool + (llama_seq_id) i, -1, -1); + } } return out; } +void engine::generate_open(llama_seq_id trunk, llama_pos pos0, const std::vector & inputs, int open_idx, + const std::vector & scored, const options & opt, result & r) { + const field_input & open = inputs[open_idx]; + // generation prefix: the scored closed fields in schema order, then the open field's suffix + std::string gen; + for (size_t f = 0; f < inputs.size(); ++f) { + if ((int) f == open_idx) { + continue; + } + const auto & fr = scored[f]; + if (fr.winner < 0 || (size_t) fr.winner >= inputs[f].candidates.size()) { + throw std::runtime_error("a closed field has no selected value"); + } + gen += inputs[f].suffix + inputs[f].candidates[fr.winner] + ",\n"; + } + gen += open.suffix; + + const tokens_t prefix = tokenize(gen, false); + const int n_batch = (int) llama_n_batch(ctx); + llama_batch batch = llama_batch_init(std::max(1, std::min(n_batch, (int) prefix.size())), 0, 1); + for (size_t i = 0; i < prefix.size(); ++i) { + common_batch_add(batch, prefix[i], pos0 + (llama_pos) i, { trunk }, i + 1 == prefix.size()); + } + const int rc = llama_decode(ctx, batch); + llama_batch_free(batch); + if (rc != 0) { + throw std::runtime_error(rc == 1 ? "no free KV cache space for the open field prefix" + : "llama_decode failed on the open field prefix (" + std::to_string(rc) + ")"); + } + std::mt19937 rng(std::random_device{}()); + std::vector scratch; + const float * logits = llama_get_logits_ith(ctx, (int) prefix.size() - 1); + llama_token next = sample_token(vocab, logits, opt.open_sampling, opt.open_temp, rng, scratch); + r.has_open = true; + for (int t = 0; t < open.max_tokens && !llama_vocab_is_eog(vocab, next); ++t) { + r.open_text += common_token_to_piece(vocab, next); + r.open_tokens += 1; + llama_batch b = llama_batch_init(1, 0, 1); + common_batch_add(b, next, pos0 + (llama_pos) (prefix.size() + t), { trunk }, true); + const int rc = llama_decode(ctx, b); + llama_batch_free(b); + if (rc != 0) { + throw std::runtime_error(rc == 1 ? "no free KV cache space for the open field generation" + : "llama_decode failed on the open field generation (" + std::to_string(rc) + ")"); + } + logits = llama_get_logits_ith(ctx, 0); + next = sample_token(vocab, logits, opt.open_sampling, opt.open_temp, rng, scratch); + } + r.open_truncated = r.open_tokens == open.max_tokens && !llama_vocab_is_eog(vocab, next); + r.open_text = recover_open_text(std::move(r.open_text), r.open_truncated, std::string("{\n") + gen, open.name); +} + // ---------------------------------------------------------------- schema compiler namespace { @@ -525,6 +674,19 @@ field_spec make_field(const std::string & name, const std::string & type, const f.name = name; f.description = description; const std::string kind = type; + if (kind == "string" && spec.contains("max_tokens")) { + if (!spec.at("max_tokens").is_number_integer()) { + throw std::invalid_argument("field \"" + name + "\": max_tokens must be an integer"); + } + const int n = spec.at("max_tokens").get(); + if (n < 1 || n > 1024) { + throw std::invalid_argument("field \"" + name + "\": open fields need max_tokens between 1 and 1024"); + } + f.type = "string"; + f.is_open = true; + f.max_tokens = n; + return f; + } if (kind == "boolean") { f.type = "boolean"; f.values = { common_json(true), common_json(false) }; @@ -633,8 +795,28 @@ compiled_schema compile_schema(const common_json & schema, const std::string & i cs.specs.push_back(make_field(e.key(), type, description, spec, json_schema)); } + int n_open = 0; + for (const auto & f : cs.specs) { + if (f.is_open) { + ++n_open; + } + } + if (n_open > 1) { + throw std::invalid_argument("at most one open (free-text) field is allowed"); + } std::string catalog; for (const auto & f : cs.specs) { + if (f.is_open) { + field_input in; + in.name = f.name; + in.suffix = " " + json_text(f.name) + ": "; + in.max_tokens = f.max_tokens; + cs.inputs.push_back(in); + catalog += (catalog.empty() ? "" : "\n") + json_text(f.name) + + (f.description.empty() ? "" : ": " + f.description) + + "\nFree text (up to " + std::to_string(f.max_tokens) + " tokens)"; + continue; + } // the value's common leading characters are fixed in the suffix; only the rest is scored std::string common = f.encoded[0]; for (const auto & v : f.encoded) { @@ -645,6 +827,7 @@ compiled_schema compile_schema(const common_json & schema, const std::string & i common.resize(c); } field_input in; + in.name = f.name; in.suffix = " " + json_text(f.name) + ": " + common; for (const auto & v : f.encoded) { in.candidates.push_back(v.substr(common.size())); @@ -658,8 +841,9 @@ compiled_schema compile_schema(const common_json & schema, const std::string & i catalog += (catalog.empty() ? "" : "\n") + json_text(f.name) + (f.description.empty() ? "" : ": " + f.description) + "\nAllowed values: " + allowed; } - cs.system_text = "Select the requested field value from its allowed values, based on the context. " - "Respond with the JSON value only.\n\nFields:\n" + catalog + "\n" + instructions; + cs.system_text = std::string(n_open ? "Select the requested field values from their allowed values and write the free-text field, based on the context. " + : "Select the requested field value from its allowed values, based on the context. ") + + "Respond with the JSON " + (n_open ? "object" : "value") + " only.\n\nFields:\n" + catalog + "\n" + instructions; return cs; } @@ -693,6 +877,16 @@ common_json assemble(const compiled_schema & cs, const result & r) { common_json fields = common_json::object(); for (size_t i = 0; i < cs.specs.size(); ++i) { const auto & sp = cs.specs[i]; + if (sp.is_open) { + common_json f = common_json::object(); + decision[sp.name] = r.open_text; + f["value"] = r.open_text; + f["generated"] = true; + f["tokens"] = r.open_tokens; + f["truncated"] = r.open_truncated; + fields[sp.name] = f; + continue; + } const auto & fr = r.fields[i]; int idx = fr.winner; common_json f = common_json::object(); diff --git a/tools/parallel-decision/decision-engine.h b/tools/parallel-decision/decision-engine.h index 49e7f104fb44..82139273bda4 100644 --- a/tools/parallel-decision/decision-engine.h +++ b/tools/parallel-decision/decision-engine.h @@ -24,9 +24,13 @@ namespace llama_decision { using tokens_t = std::vector; // One field as the scorer sees it: the text before its value and the allowed value texts. +// An open field has no candidates: its value is generated (bounded by max_tokens) on the +// context's trunk after the closed fields are scored. struct field_input { + std::string name; // schema field name std::string suffix; // e.g. ' "fire": ' std::vector candidates; // allowed values, with the suffix's shared prefix removed + int max_tokens = 0; // open fields: generation cap (1-1024) }; struct options { @@ -34,6 +38,8 @@ struct options { size_t tree_max = 128; bool split_boundary = false; // legacy: tokenise suffix and values separately bool allow_cache = true; // reuse the cached static prefix when it matches + std::string open_sampling = "greedy"; // open fields: greedy (argmax) or temperature + float open_temp = 0.7f; // open fields: temperature when open_sampling is temperature }; struct field_result { @@ -53,6 +59,12 @@ struct result { int rounds = 0; double prefill_ms = 0; double scoring_ms = 0; + // the open field (at most one): generated on the context's trunk after the closed fields + bool has_open = false; + std::string open_text; + int open_tokens = 0; + bool open_truncated = false; + double generation_ms = 0; }; // Several contexts decided against one schema and one cached prefix. Items carry fields, @@ -65,6 +77,7 @@ struct batch_result { int rounds = 0; double prefill_ms = 0; double scoring_ms = 0; + double generation_ms = 0; }; // Scores decisions on an existing context with the sequence ids [seq_base, seq_base + n_seqs): @@ -107,18 +120,26 @@ class engine { void decode_parts(const std::vector & parts); bool prepare_prefix(const tokens_t & shared, bool allow_cache); std::vector> score_branches(const std::vector & branches, llama_seq_id first, int n_free); + + // Generate the open field's value on the context's trunk (shared + context) before the trunk + // is released: the scored closed fields form the generation prefix, then the model continues + // autoregressively until end-of-generation or the field's max_tokens. + void generate_open(llama_seq_id trunk, llama_pos pos0, const std::vector & inputs, int open_idx, + const std::vector & scored, const options & opt, result & r); }; // ---- schema compiler (the C++ counterpart of llama-mojo's tools/prepare_decisions.py) struct field_spec { std::string name; - std::string type; // boolean | enum | integer | number + std::string type; // boolean | enum | integer | number | string (open) std::string description; std::string aggregate; // mode | median | mean (median/mean: numeric fields) std::vector values; // typed values; index = candidate index std::vector numbers; // numeric fields: the same values as doubles std::vector encoded; // JSON text of each value + bool is_open = false; // string field with max_tokens: generated, not scored + int max_tokens = 0; // open fields: generation cap (1-1024) }; struct compiled_schema { diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 41c6da3e1bbc..fce96a0b6776 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -2417,29 +2417,42 @@ struct server_context_impl { opt.mode = body.value("mode", std::string("auto")); opt.tree_max = (size_t) body.value("tree_max", 128); opt.allow_cache = body.value("cache_prompt", true); + opt.open_sampling = body.value("open_sampling", std::string("greedy")); + opt.open_temp = (float) body.value("open_temp", 0.7); + if (opt.open_sampling != "greedy" && opt.open_sampling != "temperature") { + throw std::invalid_argument("\"open_sampling\" must be greedy or temperature"); + } + if (opt.open_temp < 0.0f) { + throw std::invalid_argument("\"open_temp\" must be non-negative"); + } const auto b = decision_engine->decide_batch(shared, dynamic, cs.inputs, opt); size_t context_tokens = 0; + size_t generated_tokens = 0; for (const auto & r : b.items) { - context_tokens += r.context_tokens; + context_tokens += r.context_tokens; + generated_tokens += r.open_tokens; } json usage = json::object(); - usage["prompt_tokens"] = (long long) (b.shared_tokens + context_tokens); - usage["cached_tokens"] = (long long) (b.cache_hit ? b.shared_tokens : 0); - usage["context_tokens"] = (long long) context_tokens; - usage["scored_rows"] = b.rows; + usage["prompt_tokens"] = (long long) (b.shared_tokens + context_tokens); + usage["cached_tokens"] = (long long) (b.cache_hit ? b.shared_tokens : 0); + usage["context_tokens"] = (long long) context_tokens; + usage["scored_rows"] = b.rows; + usage["generated_tokens"] = (long long) generated_tokens; json timings = json::object(); timings["prefill_ms"] = b.prefill_ms; timings["scoring_ms"] = b.scoring_ms; - timings["total_ms"] = b.prefill_ms + b.scoring_ms; + timings["generation_ms"] = b.generation_ms; + timings["total_ms"] = b.prefill_ms + b.scoring_ms + b.generation_ms; timings["rounds"] = b.rounds; - timings["per_decision_ms"] = (b.prefill_ms + b.scoring_ms) / (double) b.items.size(); + timings["per_decision_ms"] = (b.prefill_ms + b.scoring_ms + b.generation_ms) / (double) b.items.size(); json results = json::array(); for (const auto & r : b.items) { json item = llama_decision::assemble(cs, r); - item["usage"] = { { "context_tokens", (long long) r.context_tokens }, { "scored_rows", r.rows } }; + item["usage"] = { { "context_tokens", (long long) r.context_tokens }, { "scored_rows", r.rows }, + { "generated_tokens", (long long) r.open_tokens } }; results.push_back(item); } json out = json::object();