Skip to content

Commit 9b22a8b

Browse files
committed
Fix parallel-call eval scoring, Muon update, and assorted correctness issues
Evaluation (training/eval.py) - Score calls as multisets so identical parallel calls are no longer collapsed by set-dedup (fixes call precision/recall on parallel). - Match predicted args to distinct reference instances via greedy bipartite matching instead of always ref[0], fixing value_acc and the param metrics for parallel_multiple. - Add BFCL-style category buckets (simple / multiple / parallel / parallel_multiple) plus irrelevance, relevance and false-trigger rates. - Split the scoring out into a pure score_tool_calls() so it can be exercised without a checkpoint. Training - Muon: apply Nesterov momentum to the raw gradient and orthogonalize the blended update. The previous order orthogonalized first and ran momentum on the orthogonal factors, so the applied update was a sum of orthogonal matrices (not itself orthogonal). Also add aspect-ratio update scaling. - Mask z-loss to non-pad positions in both train and pretrain; CE was already masked, so z-loss magnitude was scaling with the padding fraction. Inference (model/run.py, model/architecture.py) - decode() can project only the current position to the vocab, avoiding the full-buffer (B,T,d)@(d,V) matmul at every step. Numerically identical. - Stream the decode delta of the whole sequence rather than decoding each token in isolation, which dropped SentencePiece spaces and split multibyte byte-fallback pieces. Data (dataset/dataset.py, dataset/generate.py) - Weight all repeated tool names / argument values in the loss mask, not just the first occurrence (matters for parallel calls). - Match the boolean-polarity validator on whole words rather than substrings ("on" was matching inside "location", "song", ...), which had effectively disabled the inverted-boolean filter. Signed-off-by: Joe Khawand <93840910+Joe-Khawand@users.noreply.github.com>
1 parent 6fdddb8 commit 9b22a8b

8 files changed

Lines changed: 282 additions & 169 deletions

File tree

‎needle/dataset/dataset.py‎

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -71,16 +71,23 @@ def download_hf_split(split="train", repo_id=HF_DATASET_REPO):
7171

7272

7373
def _mark_json_value(s, char_w, key, value_str, weight):
74-
"""Find '"key": "value_str"' or '"key": value_str' in s, mark value chars."""
74+
"""Mark value chars for every '"key": "value_str"' / '"key": value_str' in s.
75+
76+
Marks all occurrences, not just the first, so repeated names/values across
77+
parallel calls are weighted rather than left at base weight.
78+
"""
79+
found = False
7580
pattern_str = f'"{_re.escape(key)}"\\s*:\\s*"{_re.escape(value_str)}"'
7681
for m in _re.finditer(pattern_str, s):
7782
tail = s[m.start() + len(f'"{key}"'):m.end()]
7883
val_offset = tail.index(f'"{value_str}"') + 1
7984
val_start = m.start() + len(f'"{key}"') + val_offset
8085
val_end = val_start + len(value_str)
8186
char_w[val_start:val_end] = np.maximum(char_w[val_start:val_end], weight)
87+
found = True
88+
if found:
8289
return
83-
90+
8491
pattern_ns = f'"{_re.escape(key)}"\\s*:\\s*{_re.escape(value_str)}'
8592
for m in _re.finditer(pattern_ns, s):
8693
colon_offset = s[m.start():m.end()].index(':')
@@ -89,7 +96,6 @@ def _mark_json_value(s, char_w, key, value_str, weight):
8996
val_start += 1
9097
val_end = m.end()
9198
char_w[val_start:val_end] = np.maximum(char_w[val_start:val_end], weight)
92-
return
9399

94100

95101
def _mark_json_key_in_args(s, char_w, key, weight):

‎needle/dataset/generate.py‎

Lines changed: 22 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1247,6 +1247,27 @@ def _grounding_check(pname, pval, pdesc, query, query_lower):
12471247
return True
12481248

12491249

1250+
def _boolean_enable_conflict(query_lower, pval):
1251+
"""True if a boolean 'enabled' value clearly contradicts the query intent.
1252+
1253+
Matches whole words/phrases, not substrings: "on" in "location"/"song"
1254+
otherwise makes the on-intent always fire and disables the filter.
1255+
"""
1256+
words = set(re.findall(r"[a-z']+", query_lower))
1257+
off_words = {"off", "disable", "disabled", "stop", "deactivate", "without"}
1258+
on_words = {"on", "enable", "enabled", "start", "activate"}
1259+
off_phrases = ("turn off", "switch off", "shut off", "turn it off", "don't")
1260+
on_phrases = ("turn on", "switch on", "turn it on")
1261+
wants_off = bool(words & off_words) or any(p in query_lower for p in off_phrases)
1262+
wants_on = bool(words & on_words) or any(p in query_lower for p in on_phrases)
1263+
# Only flag clear contradictions — if ambiguous, allow.
1264+
if wants_off and not wants_on and pval is True:
1265+
return True
1266+
if wants_on and not wants_off and pval is False:
1267+
return True
1268+
return False
1269+
1270+
12501271
def _semantic_check(tool_name, args, schema, query, call_type="single"):
12511272
"""Lightweight rule-based semantic validation of argument values.
12521273
@@ -1297,14 +1318,7 @@ def _semantic_check(tool_name, args, schema, query, call_type="single"):
12971318
# check alignment with query intent
12981319
if expected_type == "boolean" and isinstance(pval, bool):
12991320
if pname == "enabled":
1300-
disable_words = ("off", "disable", "stop", "don't", "no ", "without")
1301-
enable_words = ("on", "enable", "start", "turn on", "activate")
1302-
query_wants_off = any(w in query_lower for w in disable_words)
1303-
query_wants_on = any(w in query_lower for w in enable_words)
1304-
# Only reject clear contradictions — if ambiguous, allow
1305-
if query_wants_off and not query_wants_on and pval is True:
1306-
return False
1307-
if query_wants_on and not query_wants_off and pval is False:
1321+
if _boolean_enable_conflict(query_lower, pval):
13081322
return False
13091323

13101324
# Grounding check: argument values must be traceable to the query.

‎needle/model/architecture.py‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -355,11 +355,19 @@ def encode(self, src, src_mask=None):
355355
"""Backward-compatible alias for encode_text."""
356356
return self.encode_text(src, src_mask=src_mask)
357357

358-
def decode(self, tgt, encoder_out, self_mask=None, cross_mask=None, deterministic=True):
359-
"""Decode from encoder output with cross_mask for variable-length encoder output."""
358+
def decode(self, tgt, encoder_out, self_mask=None, cross_mask=None, deterministic=True, cur_pos=None):
359+
"""Decode from encoder output with cross_mask for variable-length encoder output.
360+
361+
When cur_pos (a possibly-traced scalar) is given, only that position is
362+
projected to the vocabulary — returns (B, 1, V) and skips the full-buffer
363+
(B, T, d) @ (d, V) matmul that dominates each decode step. cur_pos=None
364+
keeps the training path unchanged.
365+
"""
360366
x = self.embedding(tgt) * self.embed_scale
361367
rope = self._rope(tgt.shape[1])
362368
x = self.decoder(x, encoder_out, self_mask=self_mask, cross_mask=cross_mask, rope=rope, deterministic=deterministic)
369+
if cur_pos is not None:
370+
x = jax.lax.dynamic_slice_in_dim(x, cur_pos, 1, axis=1)
363371
logits = x.astype(jnp.float32) @ self.embedding.embedding.T
364372
return logits
365373

‎needle/model/run.py‎

Lines changed: 21 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -71,10 +71,10 @@ def _get_decode_fn(model, max_gen_len):
7171
tgt_mask = make_causal_mask(max_gen_len)
7272

7373
@jax.jit
74-
def decode_step(params, dec_buffer, encoder_out, cross_mask):
74+
def decode_step(params, dec_buffer, encoder_out, cross_mask, cur_pos):
7575
return model.apply(
7676
{"params": params}, dec_buffer, encoder_out,
77-
self_mask=tgt_mask, cross_mask=cross_mask, method="decode",
77+
self_mask=tgt_mask, cross_mask=cross_mask, cur_pos=cur_pos, method="decode",
7878
)
7979

8080
_decode_fn_cache[key] = decode_step
@@ -140,10 +140,11 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=DEFAULT_MA
140140
sys.stdout.write(f"\n")
141141
sys.stdout.flush()
142142

143-
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask)
143+
streamed = ""
144+
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask, jnp.array(0, dtype=jnp.int32))
144145

145146
for i in range(0, max_gen_len - 1):
146-
next_logits = logits[0, i]
147+
next_logits = logits[0, 0]
147148

148149
if constrained_decoder and constrained_decoder.is_active(0):
149150
logits_np = np.array(next_logits)
@@ -162,10 +163,18 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=DEFAULT_MA
162163
dec_buffer = dec_buffer.at[0, i + 1].set(next_token)
163164

164165
if stream:
165-
sys.stdout.write(tokenizer.decode([next_token]))
166-
sys.stdout.flush()
167-
168-
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask)
166+
# Emit the decode delta of the whole sequence, not the token alone:
167+
# per-token decode drops SentencePiece spaces and splits multibyte
168+
# pieces. Hold until the trailing char completes (no U+FFFD).
169+
full = tokenizer.decode(generated_tokens)
170+
if not full.endswith("�"):
171+
delta = full[len(streamed):]
172+
if delta:
173+
sys.stdout.write(delta)
174+
sys.stdout.flush()
175+
streamed = full
176+
177+
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask, jnp.array(i + 1, dtype=jnp.int32))
169178

170179
if stream:
171180
sys.stdout.write("\n")
@@ -229,18 +238,18 @@ def generate_batch(model, params, tokenizer, queries, tools_list, max_gen_len=DE
229238
from .constrained import build_constrained_decoder
230239
constrained_decoder = build_constrained_decoder(tools_list, tokenizer)
231240

232-
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask)
241+
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask, jnp.array(0, dtype=jnp.int32))
233242

234243
for pos in range(0, max_gen_len - 1):
235244
for i in range(B):
236245
if finished[i]:
237246
continue
238247
if constrained_decoder and constrained_decoder.is_active(i):
239-
logits_np = np.array(logits[i, pos])
248+
logits_np = np.array(logits[i, 0])
240249
logits_np = constrained_decoder.constrain_logits(logits_np, i)
241250
next_token = int(np.argmax(logits_np))
242251
else:
243-
next_token = int(jnp.argmax(logits[i, pos]))
252+
next_token = int(jnp.argmax(logits[i, 0]))
244253
if constrained_decoder:
245254
constrained_decoder.update(i, next_token)
246255
if next_token == eos_id:
@@ -252,7 +261,7 @@ def generate_batch(model, params, tokenizer, queries, tools_list, max_gen_len=DE
252261
if all(finished):
253262
break
254263

255-
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask)
264+
logits = decode_fn(params, dec_buffer, encoder_out, enc_mask, jnp.array(pos + 1, dtype=jnp.int32))
256265

257266
results = []
258267
for i in range(B):

0 commit comments

Comments
 (0)