From 4f30cec375daeba79df3c3236df376a74c7c5452 Mon Sep 17 00:00:00 2001 From: Karen Mosoyan Date: Tue, 10 Mar 2026 02:41:26 +0000 Subject: [PATCH 1/6] added proper shuffling --- src/cli.py | 4 +++ src/data.py | 82 +++++++++++++++++++++++++++----------------- src/test.py | 56 +++++++++++++++++++++++++----- src/tokenize_data.py | 7 +++- src/train.py | 14 +++++--- 5 files changed, 118 insertions(+), 45 deletions(-) diff --git a/src/cli.py b/src/cli.py index 15a909e..0da6b44 100644 --- a/src/cli.py +++ b/src/cli.py @@ -72,6 +72,10 @@ def main(): help="Max decoder sequence length (default: 1024)") p.add_argument("--batch-size", type=int, default=None, help="Process in batches of this size, uploading shards to GCS incrementally") + p.add_argument("--shuffle-before-split", action="store_true", + help="Shuffle the unified dataset before splitting into train/val") + p.add_argument("--split-seed", type=int, default=42, + help="Seed for --shuffle-before-split (default: 42)") p = sub.add_parser("run", add_help=False) p.add_argument("--checkpoint", type=str, required=True) diff --git a/src/data.py b/src/data.py index 5d40484..7cf5100 100644 --- a/src/data.py +++ b/src/data.py @@ -280,28 +280,22 @@ def _gcs_download_shards(cache_id, n_shards, shard_suffixes): return True -def load_tool_calls(split="train", max_samples=None, return_global_indices=False): +def load_tool_calls(split="train", max_samples=None, return_global_indices=False, + shuffle_before_split=False, shuffle_seed=42): """Load tool-calling dataset, splitting 90/10 for train/val. If return_global_indices is True, also return a numpy array mapping each split-local row position back to its row id in the full unified dataset. """ ds = _load_unified_dataset() - n = len(ds) - if split in ("validation", "val", "test"): - start, end = int(n * 0.9), n - elif split == "train": - start, end = 0, int(n * 0.9) - else: - start, end = 0, n - - global_indices = np.arange(start, end, dtype=np.int64) - ds = ds.select(range(start, end)) - - if max_samples: - limit = min(max_samples, len(ds)) - ds = ds.select(range(limit)) - global_indices = global_indices[:limit] + global_indices = _split_global_indices( + len(ds), + split=split, + max_samples=max_samples, + shuffle_before_split=shuffle_before_split, + shuffle_seed=shuffle_seed, + ) + ds = ds.select(global_indices.tolist()) if return_global_indices: return ds, global_indices @@ -329,8 +323,30 @@ def _mel_cache_key(prefix, n_samples, n_mels, max_mel_len): return hashlib.md5(key.encode()).hexdigest()[:12] +def _split_global_indices(n, split="train", max_samples=None, + shuffle_before_split=False, shuffle_seed=42): + """Return global row ids for the requested split of the unified dataset.""" + if shuffle_before_split: + indices = np.random.default_rng(shuffle_seed).permutation(n).astype(np.int64) + else: + indices = np.arange(n, dtype=np.int64) + + cut = int(n * 0.9) + if split in ("validation", "val", "test"): + indices = indices[cut:] + elif split == "train": + indices = indices[:cut] + + if max_samples: + indices = indices[:min(max_samples, len(indices))] + + return indices + + def _save_cache_metadata(split, text_cache_id, mel_cache_id, n_samples, - max_enc_len, max_dec_len, n_mels, max_mel_len): + max_enc_len, max_dec_len, n_mels, max_mel_len, + split_max_samples=None, shuffle_before_split=False, + split_seed=42): """Save metadata JSON for a split, upload to GCS.""" os.makedirs(CACHE_DIR, exist_ok=True) meta = { @@ -342,6 +358,9 @@ def _save_cache_metadata(split, text_cache_id, mel_cache_id, n_samples, "max_dec_len": max_dec_len, "n_mels": n_mels, "max_mel_len": max_mel_len, + "split_max_samples": split_max_samples, + "shuffle_before_split": shuffle_before_split, + "split_seed": split_seed, } meta_path = os.path.join(CACHE_DIR, f"{split}_metadata.json") with open(meta_path, "w") as f: @@ -610,25 +629,20 @@ def get_batches(enc_inputs, dec_inputs, dec_targets, batch_size, shuffle=True, l -def load_tool_call_audio(split="train", max_samples=None): +def load_tool_call_audio(split="train", max_samples=None, + shuffle_before_split=False, shuffle_seed=42): """Return dataset-global indices for the given split. Applies the same 90/10 split as load_tool_calls. Audio is NOT loaded into memory. """ ds = _load_unified_dataset() - n = len(ds) - - if split in ("validation", "val", "test"): - start, end = int(n * 0.9), n - elif split == "train": - start, end = 0, int(n * 0.9) - else: - start, end = 0, n - - indices = list(range(start, end)) - if max_samples: - indices = indices[:max_samples] - return indices + return _split_global_indices( + len(ds), + split=split, + max_samples=max_samples, + shuffle_before_split=shuffle_before_split, + shuffle_seed=shuffle_seed, + ).tolist() def load_audio_for_index(idx): @@ -931,6 +945,9 @@ def _shard_paths(suffix): result["kept_indices"] = np.load(cache_path + "_kept_idx.npy", mmap_mode=mmap_mode) result["mel_cache_id"] = meta.get("mel_cache_id") + result["split_max_samples"] = meta.get("split_max_samples") + result["shuffle_before_split"] = meta.get("shuffle_before_split", False) + result["split_seed"] = meta.get("split_seed", 42) return result tc_suffixes = ["_enc.npy", "_dec_in.npy", "_dec_tgt.npy", "_loss_mask.npy", "_kept_idx.npy"] @@ -948,6 +965,9 @@ def _shard_paths(suffix): "loss_mask": np.load(cache_path + "_loss_mask.npy", mmap_mode=mmap_mode), "kept_indices": np.load(cache_path + "_kept_idx.npy", mmap_mode=mmap_mode), "mel_cache_id": meta.get("mel_cache_id"), + "split_max_samples": meta.get("split_max_samples"), + "shuffle_before_split": meta.get("shuffle_before_split", False), + "split_seed": meta.get("split_seed", 42), } diff --git a/src/test.py b/src/test.py index 93502d2..23b1656 100644 --- a/src/test.py +++ b/src/test.py @@ -7,7 +7,15 @@ import numpy as np import optax -from .data import get_batches, get_tokenizer, load_tool_calls, prepare_tool_call_pairs, load_tool_call_audio, load_example_with_audio +from .data import ( + _load_cache_metadata, + get_batches, + get_tokenizer, + load_tool_calls, + prepare_tool_call_pairs, + load_tool_call_audio, + load_example_with_audio, +) from .model import ( EncoderDecoderTransformer, TransformerConfig, @@ -168,13 +176,19 @@ def compute_wer(hypotheses, references): return total_edits / max(total_ref_words, 1) -def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len=512): +def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len=512, + shuffle_before_split=False, shuffle_seed=42): """Generate tool-call predictions and compute structured metrics.""" import json from .run import generate from .data import load_tool_calls - ds = load_tool_calls("validation", max_samples=num_samples) + ds = load_tool_calls( + "validation", + max_samples=num_samples, + shuffle_before_split=shuffle_before_split, + shuffle_seed=shuffle_seed, + ) total = 0 exact_match = 0 @@ -268,12 +282,18 @@ def call_key(c): } -def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_gen_len=512): +def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_gen_len=512, + shuffle_before_split=False, shuffle_seed=42): """Generate tool-call predictions from audio and compute structured metrics.""" import json from .run import generate_from_audio - indices = load_tool_call_audio("validation", max_samples=num_samples) + indices = load_tool_call_audio( + "validation", + max_samples=num_samples, + shuffle_before_split=shuffle_before_split, + shuffle_seed=shuffle_seed, + ) total = 0 exact_match = 0 @@ -373,6 +393,9 @@ def main(args): params, config = load_checkpoint(args.checkpoint) model = EncoderDecoderTransformer(config) tokenizer = get_tokenizer() + val_meta = _load_cache_metadata("val") or {} + split_shuffle = val_meta.get("shuffle_before_split", False) + split_seed = val_meta.get("split_seed", 42) param_count = sum(x.size for x in jax.tree.leaves(params)) print(f"\ncheckpoint: {args.checkpoint}") @@ -380,7 +403,12 @@ def main(args): print(f"config: d={config.d_model}, heads={config.num_heads}, layers={config.num_encoder_layers}/{config.num_decoder_layers}") print(f"\nevaluating tool-call perplexity ({args.max_eval_samples} samples)...") - ds = load_tool_calls("validation", max_samples=args.max_eval_samples) + ds = load_tool_calls( + "validation", + max_samples=args.max_eval_samples, + shuffle_before_split=split_shuffle, + shuffle_seed=split_seed, + ) enc_inputs, dec_inputs, dec_targets, loss_mask_arr, _ = prepare_tool_call_pairs( ds, tokenizer, max_enc_len=args.max_enc_len, max_dec_len=args.max_dec_len ) @@ -390,7 +418,13 @@ def main(args): tc = None if tc_samples > 0: print(f"\nevaluating tool-call accuracy ({tc_samples} samples)...") - tc = benchmark_tool_calls(model, params, tokenizer, num_samples=tc_samples, max_gen_len=args.max_gen_len) + tc = benchmark_tool_calls( + model, params, tokenizer, + num_samples=tc_samples, + max_gen_len=args.max_gen_len, + shuffle_before_split=split_shuffle, + shuffle_seed=split_seed, + ) print(f"\n ─────────────────────────────────────") print(f" Tool-Call Metrics") @@ -420,7 +454,13 @@ def main(args): voice_tc_samples = getattr(args, "voice_tc_samples", 50) if voice_tc_samples > 0: print(f"\nevaluating voice-to-tool-call ({voice_tc_samples} samples)...") - vtc = benchmark_voice_tool_calls(model, params, tokenizer, num_samples=voice_tc_samples, max_gen_len=args.max_gen_len) + vtc = benchmark_voice_tool_calls( + model, params, tokenizer, + num_samples=voice_tc_samples, + max_gen_len=args.max_gen_len, + shuffle_before_split=split_shuffle, + shuffle_seed=split_seed, + ) print(f"\n ─── Voice-Tool-Call Metrics ─────────") print(f" JSON parse rate {vtc['json_parse_rate']:>10.1%}") print(f" Exact match {vtc['exact_match']:>10.1%}") diff --git a/src/tokenize_data.py b/src/tokenize_data.py index 04a6b5b..d0b2c6a 100644 --- a/src/tokenize_data.py +++ b/src/tokenize_data.py @@ -72,6 +72,8 @@ def tokenize(args): split=split, max_samples=args.max_samples, return_global_indices=True, + shuffle_before_split=getattr(args, "shuffle_before_split", False), + shuffle_seed=getattr(args, "split_seed", 42), ) _, _, _, _, kept_indices = prepare_tool_call_pairs( ds, tokenizer, max_enc_len=max_enc_len, max_dec_len=max_dec_len, @@ -85,7 +87,10 @@ def tokenize(args): ) _save_cache_metadata(split, text_cache_id, mel_cache_id, len(kept_indices), - max_enc_len, max_dec_len, n_mels, max_mel_len) + max_enc_len, max_dec_len, n_mels, max_mel_len, + split_max_samples=args.max_samples, + shuffle_before_split=getattr(args, "shuffle_before_split", False), + split_seed=getattr(args, "split_seed", 42)) if args.cleanup and os.path.exists(CACHE_DIR): print(f"\n=== Cleaning up {CACHE_DIR}/ ===") diff --git a/src/train.py b/src/train.py index 0b26e4a..92ddf57 100644 --- a/src/train.py +++ b/src/train.py @@ -889,14 +889,19 @@ def _tile_sp(arr): from .run import generate, generate_from_audio tp = measure_throughput(eval_model, eval_params, tokenizer, num_runs=5) - from .data import _load_unified_dataset - _ds_full = _load_unified_dataset() - _val_start = int(len(_ds_full) * 0.9) + from .data import load_tool_calls + _, val_global_indices = load_tool_calls( + "validation", + max_samples=val_data.get("split_max_samples"), + return_global_indices=True, + shuffle_before_split=val_data.get("shuffle_before_split", False), + shuffle_seed=val_data.get("split_seed", 42), + ) val_kept = val_data["kept_indices"] n_eval_samples = min(3, len(val_kept)) step = max(1, len(val_kept) // n_eval_samples) sample_indices = [val_kept[i * step] for i in range(n_eval_samples)] - eval_indices = np.array(sample_indices) + _val_start + eval_indices = val_global_indices[np.array(sample_indices)] unified_samples = [] for i, ds_idx in enumerate(eval_indices): @@ -979,4 +984,3 @@ def _tile_sp(arr): wandb.finish() print("\nTraining complete.") - From 70f1886e9db18dcaac77c35b69acfb7d3a732d90 Mon Sep 17 00:00:00 2001 From: Karen Mosoyan Date: Tue, 10 Mar 2026 07:02:40 +0000 Subject: [PATCH 2/6] added contrastive training stuf --- scripts/inspect_toucan.py | 36 +++ src/cli.py | 14 + src/data.py | 331 ++++++++++++++++++++++-- src/run.py | 49 +++- src/test.py | 13 +- src/tokenize_data.py | 16 +- src/tool_cfg.py | 526 ++++++++++++++++++++++++++++++++++++++ src/toucan.py | 149 +++++++++++ src/train.py | 199 +++++++++++--- 9 files changed, 1266 insertions(+), 67 deletions(-) create mode 100644 scripts/inspect_toucan.py create mode 100644 src/tool_cfg.py create mode 100644 src/toucan.py diff --git a/scripts/inspect_toucan.py b/scripts/inspect_toucan.py new file mode 100644 index 0000000..96a4a65 --- /dev/null +++ b/scripts/inspect_toucan.py @@ -0,0 +1,36 @@ +import argparse +import json + +from datasets import load_dataset + +from src.toucan import prepare_toucan_example + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--config", type=str, default="Kimi-K2") + parser.add_argument("--split", type=str, default="train") + parser.add_argument("--samples", type=int, default=5) + return parser.parse_args() + + +def main(): + args = parse_args() + ds = load_dataset("Agent-Ark/Toucan-1.5M", args.config, split=args.split, streaming=True) + for i, row in enumerate(ds): + if i >= args.samples: + break + ex = prepare_toucan_example(row) + print(f"ROW {i}") + print(f"subset_name: {ex['subset_name']}") + print(f"question: {ex['question'][:160]}") + print(f"target_tools: {ex['target_tools']}") + print(f"positive_indices: {ex['positive_indices']}") + print(f"tool_names: {ex['tool_names'][:5]}") + print(f"tools_json: {ex['tools_json'][:200]}") + print(f"tool_text[0]: {ex['tool_texts'][0][:300]}") + print() + + +if __name__ == "__main__": + main() diff --git a/src/cli.py b/src/cli.py index 0da6b44..7a6cbcb 100644 --- a/src/cli.py +++ b/src/cli.py @@ -56,6 +56,14 @@ def main(): help="Number of mel frequency bins (default: 80)") p.add_argument("--max-speech-samples", type=int, default=None, help="Max voice-tool-call training samples (default: all)") + p.add_argument("--cfg-inference", action="store_true", + help="Enable CFG-constrained sample generation during training") + p.add_argument("--tool-contrastive-weight", type=float, default=0.0, + help="Auxiliary Toucan tool-alignment loss weight") + p.add_argument("--audio-text-contrastive-weight", type=float, default=0.0, + help="Auxiliary paired audio-text SigLIP loss weight") + p.add_argument("--skip-epoch-extras", action="store_true", + help="Skip epoch-end throughput benchmark and qualitative sample generation") p = sub.add_parser("tokenize", add_help=False) p.add_argument("--max-samples", type=int, default=None, @@ -76,6 +84,10 @@ def main(): help="Shuffle the unified dataset before splitting into train/val") p.add_argument("--split-seed", type=int, default=42, help="Seed for --shuffle-before-split (default: 42)") + p.add_argument("--toucan-config", type=str, default=None, + help="Optional Toucan subset to parse and cache during tokenization") + p.add_argument("--toucan-max-samples", type=int, default=None, + help="Optional max samples for Toucan parsing") p = sub.add_parser("run", add_help=False) p.add_argument("--checkpoint", type=str, required=True) @@ -84,6 +96,7 @@ def main(): p.add_argument("--audio", type=str, nargs="*", help="Audio file paths for voice-to-tool-call") p.add_argument("--max-len", type=int, default=512) p.add_argument("--seed", type=int, default=0) + p.add_argument("--cfg-inference", action="store_true") p = sub.add_parser("test", add_help=False) p.add_argument("--checkpoint", type=str, required=True) @@ -97,6 +110,7 @@ def main(): p.add_argument("--voice-tc-samples", type=int, default=50, help="Samples for voice-to-tool-call eval (default: 50)") p.add_argument("--throughput-runs", type=int, default=10) + p.add_argument("--cfg-inference", action="store_true") p = sub.add_parser("evaluate", add_help=False) p.add_argument("--checkpoint", type=str, required=True) diff --git a/src/data.py b/src/data.py index 7cf5100..aacf9aa 100644 --- a/src/data.py +++ b/src/data.py @@ -12,6 +12,7 @@ logging.getLogger("transformers").setLevel(logging.ERROR) logging.getLogger("huggingface_hub").setLevel(logging.ERROR) +import math import numpy as np from datasets import Audio, load_from_disk from tqdm import tqdm @@ -30,6 +31,43 @@ TOOL_CALL_ID = 4 TRANSCRIBE_ID = 5 +JSON_LBRACK = "" +JSON_RBRACK = "" +JSON_LBRACE = "" +JSON_RBRACE = "" +JSON_COLON = "" +JSON_COMMA = "" +JSON_QUOTE = "" +JSON_TRUE = "" +JSON_FALSE = "" +JSON_NULL = "" +JSON_KEY_NAME = "" +JSON_KEY_PARAMETERS = "" +JSON_KEY_ARGUMENTS = "" + +JSON_SPECIAL_LITERALS = { + JSON_LBRACK: "[", + JSON_RBRACK: "]", + JSON_LBRACE: "{", + JSON_RBRACE: "}", + JSON_COLON: ":", + JSON_COMMA: ",", + JSON_QUOTE: '"', + JSON_TRUE: "true", + JSON_FALSE: "false", + JSON_NULL: "null", + JSON_KEY_NAME: '"name"', + JSON_KEY_PARAMETERS: '"parameters"', + JSON_KEY_ARGUMENTS: '"arguments"', +} +JSON_SPECIAL_SYMBOLS = tuple(JSON_SPECIAL_LITERALS.keys()) +JSON_KEY_SYMBOLS = { + "name": JSON_KEY_NAME, + "parameters": JSON_KEY_PARAMETERS, + "arguments": JSON_KEY_ARGUMENTS, +} +TRAINABLE_SPECIAL_SYMBOLS = ["", "", *JSON_SPECIAL_SYMBOLS] + _unified_dataset_cache = None @@ -79,12 +117,198 @@ def _set_audio_backend(ds): return ds +def _normalize_tool_schema_spec(tools_text): + try: + tools = _json.loads(tools_text) + except (_json.JSONDecodeError, TypeError): + return [] + + if not isinstance(tools, list): + return [] + + normalized = [] + for tool in tools: + if not isinstance(tool, dict): + continue + name = str(tool.get("name", "")) + params = tool.get("parameters") + if not isinstance(params, dict): + params = {} + ordered_params = [(str(k), params[k]) for k in params.keys()] + normalized.append({"name": name, "parameters": ordered_params}) + return normalized + + +def _normalize_tool_call_spec(answer_text, tools_text): + schema = _normalize_tool_schema_spec(tools_text) + param_order = {tool["name"]: [name for name, _ in tool["parameters"]] for tool in schema} + + try: + calls = _json.loads(answer_text) + except (_json.JSONDecodeError, TypeError): + calls = [] + + if not isinstance(calls, list): + calls = [] + + normalized = [] + for call in calls: + if not isinstance(call, dict): + continue + name = str(call.get("name", "")) + args = call.get("arguments") + if not isinstance(args, dict): + args = {} + ordered = [] + seen = set() + for arg_name in param_order.get(name, ()): + if arg_name in args: + ordered.append((arg_name, args[arg_name])) + seen.add(arg_name) + for arg_name, arg_value in args.items(): + if arg_name not in seen: + ordered.append((str(arg_name), arg_value)) + normalized.append({"name": name, "arguments": ordered}) + return normalized + + +def _escape_json_string_content(text): + content = _json.dumps(str(text), ensure_ascii=True)[1:-1] + content = content.replace('\\"', '\\u0022') + content = content.replace("<", "\\u003c").replace(">", "\\u003e") + return content + + +def _number_to_text(value): + if isinstance(value, bool): + return "true" if value else "false" + if value is None: + return "null" + if isinstance(value, int): + return str(value) + if isinstance(value, float): + if not math.isfinite(value): + return "0" + text = format(value, ".15g") + if text == "-0": + return "0" + return text + return _json.dumps(value, ensure_ascii=True, separators=(",", ":"), allow_nan=False) + + +def _string_segments(text): + return [JSON_QUOTE, _escape_json_string_content(text), JSON_QUOTE] + + +def _generic_json_segments(value): + if isinstance(value, str): + return _string_segments(value) + if value is True: + return [JSON_TRUE] + if value is False: + return [JSON_FALSE] + if value is None: + return [JSON_NULL] + if isinstance(value, (int, float)): + return [_number_to_text(value)] + if isinstance(value, list): + segments = [JSON_LBRACK] + for i, item in enumerate(value): + if i: + segments.append(JSON_COMMA) + segments.extend(_generic_json_segments(item)) + segments.append(JSON_RBRACK) + return segments + if isinstance(value, dict): + segments = [JSON_LBRACE] + for i, (key, item) in enumerate(value.items()): + if i: + segments.append(JSON_COMMA) + segments.extend(_string_segments(str(key))) + segments.append(JSON_COLON) + segments.extend(_generic_json_segments(item)) + segments.append(JSON_RBRACE) + return segments + return _string_segments(str(value)) + + +def _tool_schema_to_segments(tools_text): + schema = _normalize_tool_schema_spec(tools_text) + segments = [JSON_LBRACK] + for i, tool in enumerate(schema): + if i: + segments.append(JSON_COMMA) + segments.append(JSON_LBRACE) + segments.extend([JSON_KEY_NAME, JSON_COLON, *_string_segments(tool["name"])]) + segments.append(JSON_COMMA) + segments.extend([JSON_KEY_PARAMETERS, JSON_COLON, JSON_LBRACE]) + for j, (param_name, param_type) in enumerate(tool["parameters"]): + if j: + segments.append(JSON_COMMA) + segments.extend(_string_segments(param_name)) + segments.append(JSON_COLON) + segments.extend(_generic_json_segments(param_type)) + segments.extend([JSON_RBRACE, JSON_RBRACE]) + segments.append(JSON_RBRACK) + return segments + + +def _tool_call_to_segments(answer_text, tools_text): + calls = _normalize_tool_call_spec(answer_text, tools_text) + segments = [JSON_LBRACK] + for i, call in enumerate(calls): + if i: + segments.append(JSON_COMMA) + segments.append(JSON_LBRACE) + segments.extend([JSON_KEY_NAME, JSON_COLON, *_string_segments(call["name"])]) + segments.append(JSON_COMMA) + segments.extend([JSON_KEY_ARGUMENTS, JSON_COLON, JSON_LBRACE]) + for j, (arg_name, arg_value) in enumerate(call["arguments"]): + if j: + segments.append(JSON_COMMA) + segments.extend(_string_segments(arg_name)) + segments.append(JSON_COLON) + segments.extend(_generic_json_segments(arg_value)) + segments.extend([JSON_RBRACE, JSON_RBRACE]) + segments.append(JSON_RBRACK) + return segments + + +def _segments_to_training_text(segments): + return " ".join(segments) + + class NeedleTokenizer: """Wrapper around SentencePiece providing the interface the codebase expects.""" def __init__(self, model_path): self.sp = spm.SentencePieceProcessor() self.sp.Load(model_path) + self.json_token_ids = {} + for symbol in JSON_SPECIAL_SYMBOLS: + piece_id = int(self.sp.PieceToId(symbol)) + if piece_id < 0 or self.sp.IdToPiece(piece_id) != symbol: + raise ValueError( + "Tokenizer is missing structured JSON symbols. " + "Re-run `needle tokenize` to retrain the tokenizer." + ) + self.json_token_ids[symbol] = piece_id + self._special_id_to_literal = { + self.json_token_ids[symbol]: literal + for symbol, literal in JSON_SPECIAL_LITERALS.items() + } + excluded = { + PAD_ID, + EOS_ID, + BOS_ID, + TOOL_CALL_ID, + TRANSCRIBE_ID, + *self._special_id_to_literal.keys(), + } + self.regular_token_ids = np.array( + [i for i in range(self.sp.GetPieceSize()) if i not in excluded], + dtype=np.int32, + ) @property def pad_token_id(self): @@ -113,10 +337,59 @@ def vocab_size(self): def encode(self, text): return self.sp.Encode(text, out_type=int) + def _encode_segments(self, segments): + ids = [] + for segment in segments: + token_id = self.json_token_ids.get(segment) + if token_id is not None: + ids.append(token_id) + elif segment: + ids.extend(self.sp.Encode(segment, out_type=int)) + return ids + + def encode_tool_schema(self, tools_text): + return self._encode_segments(_tool_schema_to_segments(tools_text)) + + def encode_tool_call(self, answer_text, tools_text): + return self._encode_segments(_tool_call_to_segments(answer_text, tools_text)) + + def encode_json_string_content(self, text): + return self.sp.Encode(_escape_json_string_content(text), out_type=int) + + def encode_json_number(self, value): + return self.sp.Encode(_number_to_text(value), out_type=int) + + def decode_structured(self, ids): + if isinstance(ids, (list, tuple)) and len(ids) > 0 and isinstance(ids[0], (list, tuple, np.ndarray)): + return [self.decode_structured(seq) for seq in ids] + + out = [] + regular = [] + for token_id in list(ids): + token_id = int(token_id) + literal = self._special_id_to_literal.get(token_id) + if literal is None: + regular.append(token_id) + continue + if regular: + out.append(self.sp.Decode(regular)) + regular = [] + out.append(literal) + if regular: + out.append(self.sp.Decode(regular)) + return "".join(out) + def decode(self, ids): - if isinstance(ids, (list, tuple)) and len(ids) > 0 and isinstance(ids[0], (list, tuple)): - return [self.sp.Decode(seq) for seq in ids] - return self.sp.Decode(list(ids)) + return self.decode_structured(ids) + + def token_surface(self, token_id): + token_id = int(token_id) + if token_id in self._special_id_to_literal: + return self._special_id_to_literal[token_id] + return self.sp.Decode([token_id]) + + def quote_token_id(self): + return self.json_token_ids[JSON_QUOTE] def __call__(self, texts, truncation=True, max_length=None, **kwargs): all_ids = [] @@ -163,10 +436,17 @@ def train_tokenizer(vocab_size=8192, max_samples=None, force=False): corpus_path = os.path.join(TOKENIZER_DIR, "corpus.txt") with open(corpus_path, "w") as f: for example in tqdm(ds, desc="Writing corpus"): - for field in ("query", "tools", "answers"): - text = example[field].strip() - if text: - f.write(text + "\n") + query = example["query"].strip() + if query: + f.write(query + "\n") + + tools_text = _segments_to_training_text(_tool_schema_to_segments(example["tools"])) + if tools_text: + f.write(tools_text + "\n") + + answers_text = _segments_to_training_text(_tool_call_to_segments(example["answers"], example["tools"])) + if answers_text: + f.write(answers_text + "\n") spm.SentencePieceTrainer.Train( input=corpus_path, @@ -177,7 +457,7 @@ def train_tokenizer(vocab_size=8192, max_samples=None, force=False): eos_id=EOS_ID, bos_id=BOS_ID, unk_id=UNK_ID, - user_defined_symbols=["", ""], + user_defined_symbols=TRAINABLE_SPECIAL_SYMBOLS, byte_fallback=True, normalization_rule_name="identity", num_threads=os.cpu_count(), @@ -346,7 +626,7 @@ def _split_global_indices(n, split="train", max_samples=None, def _save_cache_metadata(split, text_cache_id, mel_cache_id, n_samples, max_enc_len, max_dec_len, n_mels, max_mel_len, split_max_samples=None, shuffle_before_split=False, - split_seed=42): + split_seed=42, toucan_cache_path=None): """Save metadata JSON for a split, upload to GCS.""" os.makedirs(CACHE_DIR, exist_ok=True) meta = { @@ -361,6 +641,7 @@ def _save_cache_metadata(split, text_cache_id, mel_cache_id, n_samples, "split_max_samples": split_max_samples, "shuffle_before_split": shuffle_before_split, "split_seed": split_seed, + "toucan_cache_path": toucan_cache_path, } meta_path = os.path.join(CACHE_DIR, f"{split}_metadata.json") with open(meta_path, "w") as f: @@ -457,19 +738,17 @@ def _load_tc_cache(): enc_results = pool.map(_tokenize_chunk, enc_chunks) all_enc_tokens = [tok for chunk in enc_results for tok in chunk] - tools_chunks = [tools_texts[i:i + chunk_size] for i in range(0, len(tools_texts), chunk_size)] - print(f"Tokenizing tools ({num_workers} workers)...") - with mp.Pool(num_workers, initializer=_init_worker, - initargs=(model_path, max_dec_len - 2)) as pool: - tools_results = pool.map(_tokenize_chunk, tools_chunks) - all_tools_tokens = [tok for chunk in tools_results for tok in chunk] + print("Tokenizing tools (structured JSON)...") + all_tools_tokens = [ + tokenizer.encode_tool_schema(text)[:max_dec_len - 2] + for text in tqdm(tools_texts, desc="Encoding tools") + ] - ans_chunks = [ans_texts[i:i + chunk_size] for i in range(0, len(ans_texts), chunk_size)] - print(f"Tokenizing answers ({num_workers} workers)...") - with mp.Pool(num_workers, initializer=_init_worker, - initargs=(model_path, max_dec_len)) as pool: - ans_results = pool.map(_tokenize_chunk, ans_chunks) - all_ans_tokens = [tok for chunk in ans_results for tok in chunk] + print("Tokenizing answers (structured JSON)...") + all_ans_tokens = [ + tokenizer.encode_tool_call(answer, tools)[:max_dec_len] + for answer, tools in tqdm(zip(ans_texts, tools_texts), total=len(ans_texts), desc="Encoding answers") + ] n = len(ds) @@ -628,6 +907,14 @@ def get_batches(enc_inputs, dec_inputs, dec_targets, batch_size, shuffle=True, l yield batch +def get_text_mel_batches(enc_inputs, mel_data, batch_size, shuffle=True): + n = len(enc_inputs) + indices = np.random.permutation(n) if shuffle else np.arange(n) + for i in range(0, n - batch_size + 1, batch_size): + idx = indices[i : i + batch_size] + yield np.array(enc_inputs[idx]), np.array(mel_data[idx]) + + def load_tool_call_audio(split="train", max_samples=None, shuffle_before_split=False, shuffle_seed=42): @@ -948,6 +1235,7 @@ def _shard_paths(suffix): result["split_max_samples"] = meta.get("split_max_samples") result["shuffle_before_split"] = meta.get("shuffle_before_split", False) result["split_seed"] = meta.get("split_seed", 42) + result["toucan_cache_path"] = meta.get("toucan_cache_path") return result tc_suffixes = ["_enc.npy", "_dec_in.npy", "_dec_tgt.npy", "_loss_mask.npy", "_kept_idx.npy"] @@ -968,6 +1256,7 @@ def _shard_paths(suffix): "split_max_samples": meta.get("split_max_samples"), "shuffle_before_split": meta.get("shuffle_before_split", False), "split_seed": meta.get("split_seed", 42), + "toucan_cache_path": meta.get("toucan_cache_path"), } diff --git a/src/run.py b/src/run.py index eb3cb9b..72b942d 100644 --- a/src/run.py +++ b/src/run.py @@ -14,6 +14,7 @@ make_padding_mask, make_mel_padding_mask, ) +from .tool_cfg import ToolCallCFG _decode_fn_cache = {} @@ -48,7 +49,17 @@ def load_checkpoint(path): return params, config -def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed=0, stream=True, task_token_id=None): +def _select_next_token(next_logits, allowed_ids=None): + if allowed_ids is None: + return int(jnp.argmax(next_logits)) + allowed_ids = jnp.array(allowed_ids, dtype=jnp.int32) + masked = jnp.full_like(next_logits, jnp.finfo(next_logits.dtype).min) + masked = masked.at[allowed_ids].set(next_logits[allowed_ids]) + return int(jnp.argmax(masked)) + + +def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed=0, + stream=True, task_token_id=None, use_cfg=False): """Generate tool-call output. Encoder: query only. @@ -66,7 +77,7 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= {"params": params}, enc_input, src_mask=src_mask, method="encode" ) - tools_tokens = tokenizer.encode(tools) + tools_tokens = tokenizer.encode_tool_schema(tools) prefix = [eos_id, tool_call_id] + tools_tokens prefix_len = min(len(prefix), max_gen_len) @@ -77,6 +88,8 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= decode_fn = _get_decode_fn(model, max_gen_len) generated_tokens = [] + cfg = ToolCallCFG(tokenizer, tools) if use_cfg else None + streamed_text = "" if stream: sys.stdout.write(f"\n") @@ -86,16 +99,21 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= for i in range(prefix_len - 1, max_gen_len - 1): next_logits = logits[0, i] - next_token = int(jnp.argmax(next_logits)) + allowed_ids = cfg.allowed_next_ids(training=False) if cfg is not None else None + next_token = _select_next_token(next_logits, allowed_ids) if next_token == eos_id: break generated_tokens.append(next_token) + if cfg is not None: + cfg.step(next_token) dec_buffer = dec_buffer.at[0, i + 1].set(next_token) if stream: - sys.stdout.write(tokenizer.decode([next_token])) + current_text = tokenizer.decode_structured(generated_tokens) + sys.stdout.write(current_text[len(streamed_text):]) + streamed_text = current_text sys.stdout.flush() logits = decode_fn(params, dec_buffer, encoder_out) @@ -103,7 +121,7 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= if stream: sys.stdout.write("\n") - return tokenizer.decode(generated_tokens) + return tokenizer.decode_structured(generated_tokens) def load_audio(path, target_sr=16000): @@ -120,7 +138,8 @@ def load_audio(path, target_sr=16000): return audio, sr -def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools="[]", max_gen_len=512, seed=0, stream=True): +def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools="[]", + max_gen_len=512, seed=0, stream=True, use_cfg=False): """Generate tool-call output from audio using the speech encoder pathway. mel -> encode_speech -> decoder [BOS, , tools_tokens...] -> greedy decode. @@ -139,7 +158,7 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" {"params": params}, mel_input, src_mask=src_mask, deterministic=True, method="encode_speech" ) - tools_tokens = tokenizer.encode(tools) + tools_tokens = tokenizer.encode_tool_schema(tools) prefix = [eos_id, tool_call_id] + tools_tokens prefix_len = min(len(prefix), max_gen_len) @@ -150,6 +169,8 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" decode_fn = _get_decode_fn(model, max_gen_len) generated_tokens = [] + cfg = ToolCallCFG(tokenizer, tools) if use_cfg else None + streamed_text = "" if stream: sys.stdout.write("\n") @@ -159,16 +180,21 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" for i in range(prefix_len - 1, max_gen_len - 1): next_logits = logits[0, i] - next_token = int(jnp.argmax(next_logits)) + allowed_ids = cfg.allowed_next_ids(training=False) if cfg is not None else None + next_token = _select_next_token(next_logits, allowed_ids) if next_token == eos_id: break generated_tokens.append(next_token) + if cfg is not None: + cfg.step(next_token) dec_buffer = dec_buffer.at[0, i + 1].set(next_token) if stream: - sys.stdout.write(tokenizer.decode([next_token])) + current_text = tokenizer.decode_structured(generated_tokens) + sys.stdout.write(current_text[len(streamed_text):]) + streamed_text = current_text sys.stdout.flush() logits = decode_fn(params, dec_buffer, encoder_out) @@ -176,7 +202,7 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" if stream: sys.stdout.write("\n") - return tokenizer.decode(generated_tokens) + return tokenizer.decode_structured(generated_tokens) def main(args): @@ -206,6 +232,7 @@ def main(args): max_gen_len=args.max_len, seed=args.seed + i, stream=True, + use_cfg=getattr(args, "cfg_inference", False), ) return @@ -233,6 +260,7 @@ def main(args): max_gen_len=args.max_len, seed=args.seed + i, stream=True, + use_cfg=getattr(args, "cfg_inference", False), ) @@ -244,6 +272,7 @@ def parse_args(): parser.add_argument("--audio", type=str, nargs="*", help="Audio file paths for voice-to-tool-call") parser.add_argument("--max-len", type=int, default=512) parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--cfg-inference", action="store_true") return parser.parse_args() diff --git a/src/test.py b/src/test.py index 23b1656..857157e 100644 --- a/src/test.py +++ b/src/test.py @@ -131,7 +131,7 @@ def benchmark_generation_quality(model, params, tokenizer, prompts, max_gen_len= generations = [] for i, prompt in enumerate(prompts): - text = generate(model, params, tokenizer, prompt, max_gen_len, temperature, seed=i, stream=False) + text = generate(model, params, tokenizer, prompt, max_gen_len=max_gen_len, seed=i, stream=False) generations.append(text) lengths = [len(tokenizer.encode(t)) for t in generations] @@ -177,7 +177,7 @@ def compute_wer(hypotheses, references): def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len=512, - shuffle_before_split=False, shuffle_seed=42): + shuffle_before_split=False, shuffle_seed=42, use_cfg=False): """Generate tool-call predictions and compute structured metrics.""" import json from .run import generate @@ -209,7 +209,7 @@ def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len= pred_text = generate( model, params, tokenizer, ex["query"], - tools=ex["tools"], max_gen_len=max_gen_len, seed=i, stream=False, + tools=ex["tools"], max_gen_len=max_gen_len, seed=i, stream=False, use_cfg=use_cfg, ).strip() try: @@ -283,7 +283,7 @@ def call_key(c): def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_gen_len=512, - shuffle_before_split=False, shuffle_seed=42): + shuffle_before_split=False, shuffle_seed=42, use_cfg=False): """Generate tool-call predictions from audio and compute structured metrics.""" import json from .run import generate_from_audio @@ -316,7 +316,7 @@ def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_ge pred_text = generate_from_audio( model, params, tokenizer, pair["audio_array"], sr=pair["sampling_rate"], - tools=pair["tools"], max_gen_len=max_gen_len, seed=i, stream=False, + tools=pair["tools"], max_gen_len=max_gen_len, seed=i, stream=False, use_cfg=use_cfg, ).strip() try: @@ -424,6 +424,7 @@ def main(args): max_gen_len=args.max_gen_len, shuffle_before_split=split_shuffle, shuffle_seed=split_seed, + use_cfg=getattr(args, "cfg_inference", False), ) print(f"\n ─────────────────────────────────────") @@ -460,6 +461,7 @@ def main(args): max_gen_len=args.max_gen_len, shuffle_before_split=split_shuffle, shuffle_seed=split_seed, + use_cfg=getattr(args, "cfg_inference", False), ) print(f"\n ─── Voice-Tool-Call Metrics ─────────") print(f" JSON parse rate {vtc['json_parse_rate']:>10.1%}") @@ -502,6 +504,7 @@ def parse_args(): help="Samples for tool-call accuracy eval (default: 200)") parser.add_argument("--voice-tc-samples", type=int, default=50, help="Samples for voice-to-tool-call eval (default: 50)") + parser.add_argument("--cfg-inference", action="store_true") return parser.parse_args() diff --git a/src/tokenize_data.py b/src/tokenize_data.py index d0b2c6a..a20cd0c 100644 --- a/src/tokenize_data.py +++ b/src/tokenize_data.py @@ -28,6 +28,7 @@ train_tokenizer, upload_tokenizer_to_gcs, ) +from .toucan import cache_toucan_examples def _clear_gcs_caches(): @@ -58,6 +59,18 @@ def tokenize(args): upload_tokenizer_to_gcs() tokenizer = get_tokenizer() + toucan_path = None + + if getattr(args, "toucan_config", None): + print("\n=== Caching Toucan tool definitions ===") + toucan_path = cache_toucan_examples( + config=args.toucan_config, + split="train", + max_samples=getattr(args, "toucan_max_samples", None), + tokenizer=tokenizer, + max_text_len=max(getattr(args, "max_enc_len", 256), getattr(args, "max_dec_len", 1024)), + ) + print(f"Cached Toucan examples to {toucan_path}") print("\n=== Tokenizing text data + precomputing mels ===") max_enc_len = getattr(args, "max_enc_len", 256) @@ -90,7 +103,8 @@ def tokenize(args): max_enc_len, max_dec_len, n_mels, max_mel_len, split_max_samples=args.max_samples, shuffle_before_split=getattr(args, "shuffle_before_split", False), - split_seed=getattr(args, "split_seed", 42)) + split_seed=getattr(args, "split_seed", 42), + toucan_cache_path=toucan_path) if args.cleanup and os.path.exists(CACHE_DIR): print(f"\n=== Cleaning up {CACHE_DIR}/ ===") diff --git a/src/tool_cfg.py b/src/tool_cfg.py new file mode 100644 index 0000000..be68f6b --- /dev/null +++ b/src/tool_cfg.py @@ -0,0 +1,526 @@ +import re + +import numpy as np + +from .data import ( + EOS_ID, + JSON_COLON, + JSON_COMMA, + JSON_FALSE, + JSON_KEY_ARGUMENTS, + JSON_KEY_NAME, + JSON_LBRACE, + JSON_LBRACK, + JSON_NULL, + JSON_QUOTE, + JSON_RBRACE, + JSON_RBRACK, + JSON_TRUE, + _normalize_tool_schema_spec, +) + + +_NUMBER_PREFIX_RE = re.compile(r"^-?(?:0|[1-9]\d*)(?:\.\d*)?(?:[eE][+-]?\d*)?$|^-?(?:0|[1-9]\d*)?$|^-?(?:0|[1-9]\d*)\.$|^-?(?:0|[1-9]\d*)(?:\.\d+)?[eE]?$|^-?(?:0|[1-9]\d*)(?:\.\d+)?[eE][+-]?$") +_NUMBER_COMPLETE_RE = re.compile(r"^-?(?:0|[1-9]\d*)(?:\.\d+)?(?:[eE][+-]?\d+)?$") + + +def _is_number_prefix(text): + return bool(text) and _NUMBER_PREFIX_RE.match(text) is not None + + +def _is_complete_number(text): + return bool(text) and _NUMBER_COMPLETE_RE.match(text) is not None + + +def _schema_type_name(value): + if isinstance(value, str): + lowered = value.lower() + if lowered in {"str", "string", "text"}: + return "string" + if lowered in {"bool", "boolean"}: + return "boolean" + if lowered in {"int", "integer", "long"}: + return "integer" + if lowered in {"float", "double", "number", "numeric"}: + return "number" + if lowered in {"list", "array"}: + return "array" + if lowered in {"dict", "map", "object"}: + return "object" + if isinstance(value, list): + return "array" + if isinstance(value, dict): + return "object" + return "any" + + +def _build_trie(sequences): + root = {} + for index, sequence in sequences: + node = root + for token_id in sequence: + node = node.setdefault(int(token_id), {}) + node["_end"] = index + return root + + +class JsonValueParser: + def __init__(self, tokenizer, type_hint="any"): + self.tokenizer = tokenizer + self.type_hint = type_hint + self.regular_ids = tokenizer.regular_token_ids + self.quote_id = tokenizer.json_token_ids[JSON_QUOTE] + self.lbrace_id = tokenizer.json_token_ids[JSON_LBRACE] + self.rbrace_id = tokenizer.json_token_ids[JSON_RBRACE] + self.lbrack_id = tokenizer.json_token_ids[JSON_LBRACK] + self.rbrack_id = tokenizer.json_token_ids[JSON_RBRACK] + self.colon_id = tokenizer.json_token_ids[JSON_COLON] + self.comma_id = tokenizer.json_token_ids[JSON_COMMA] + self.true_id = tokenizer.json_token_ids[JSON_TRUE] + self.false_id = tokenizer.json_token_ids[JSON_FALSE] + self.null_id = tokenizer.json_token_ids[JSON_NULL] + self.mode = "expect_value" + self.stack = [] + self.number_text = "" + self.complete = False + self._regular_surfaces = { + int(token_id): tokenizer.token_surface(int(token_id)) + for token_id in self.regular_ids + } + self._number_prefix_ids = np.array( + [token_id for token_id, text in self._regular_surfaces.items() if _is_number_prefix(text)], + dtype=np.int32, + ) + self._number_cache = {} + + def _start_value_ids(self, type_hint=None): + hint = type_hint or self.type_hint + if hint == "string": + return np.array([self.quote_id], dtype=np.int32) + if hint == "boolean": + return np.array([self.true_id, self.false_id], dtype=np.int32) + if hint in {"integer", "number"}: + return self._number_prefix_ids + if hint == "array": + return np.array([self.lbrack_id], dtype=np.int32) + if hint == "object": + return np.array([self.lbrace_id], dtype=np.int32) + return np.concatenate( + [ + np.array([self.quote_id, self.lbrace_id, self.lbrack_id, self.true_id, self.false_id, self.null_id], dtype=np.int32), + self._number_prefix_ids, + ] + ) + + def _number_allowed_ids(self, training=False): + if training: + return None + cached = self._number_cache.get(self.number_text) + if cached is not None: + return cached + allowed = [] + for token_id, surface in self._regular_surfaces.items(): + if _is_number_prefix(self.number_text + surface): + allowed.append(token_id) + if self._is_value_complete(): + allowed.extend(self._value_end_ids().tolist()) + cached = np.array(sorted(set(allowed)), dtype=np.int32) + self._number_cache[self.number_text] = cached + return cached + + def _value_end_ids(self): + if not self.stack: + return np.array([EOS_ID], dtype=np.int32) + frame = self.stack[-1] + if frame["kind"] == "array": + return np.array([self.comma_id, self.rbrack_id], dtype=np.int32) + return np.array([self.comma_id, self.rbrace_id], dtype=np.int32) + + def _value_finished(self): + if not self.stack: + self.complete = True + self.mode = "done" + return + frame = self.stack[-1] + if frame["kind"] == "array": + frame["mode"] = "after_value" + self.mode = "array_after_value" + else: + frame["mode"] = "after_value" + self.mode = "object_after_value" + + def allowed_next_ids(self, training=False): + if self.mode == "done": + return np.array([EOS_ID], dtype=np.int32) + if self.mode == "expect_value": + ids = self._start_value_ids() + if training and self.type_hint == "any": + return np.array([self.quote_id, self.lbrace_id, self.lbrack_id, self.true_id, self.false_id, self.null_id], dtype=np.int32) + return ids + if self.mode == "string": + if training: + return None + return np.concatenate([self.regular_ids, np.array([self.quote_id], dtype=np.int32)]) + if self.mode == "number": + return self._number_allowed_ids(training=training) + if self.mode == "array_value_or_end": + return np.concatenate([np.array([self.rbrack_id], dtype=np.int32), self._start_value_ids("any")]) + if self.mode == "array_after_value": + return np.array([self.comma_id, self.rbrack_id], dtype=np.int32) + if self.mode == "object_key_or_end": + return np.array([self.quote_id, self.rbrace_id], dtype=np.int32) + if self.mode == "object_key_string": + if training: + return None + return np.concatenate([self.regular_ids, np.array([self.quote_id], dtype=np.int32)]) + if self.mode == "object_after_key": + return np.array([self.colon_id], dtype=np.int32) + if self.mode == "object_after_value": + return np.array([self.comma_id, self.rbrace_id], dtype=np.int32) + return None + + def step(self, token_id): + token_id = int(token_id) + if self.mode == "expect_value": + if token_id == self.quote_id: + self.mode = "string" + return + if token_id == self.lbrack_id: + self.stack.append({"kind": "array", "mode": "value_or_end"}) + self.mode = "array_value_or_end" + return + if token_id == self.lbrace_id: + self.stack.append({"kind": "object", "mode": "key_or_end"}) + self.mode = "object_key_or_end" + return + if token_id in {self.true_id, self.false_id, self.null_id}: + self._value_finished() + return + self.mode = "number" + self.number_text = self._regular_surfaces.get(token_id, "") + return + + if self.mode == "string": + if token_id == self.quote_id: + self._value_finished() + return + + if self.mode == "number": + if token_id in set(self._value_end_ids().tolist()): + if token_id == EOS_ID: + self.complete = True + self.mode = "done" + return + frame = self.stack[-1] + if frame["kind"] == "array": + if token_id == self.comma_id: + frame["mode"] = "value" + self.mode = "expect_value" + self.type_hint = "any" + else: + self.stack.pop() + self._value_finished() + else: + if token_id == self.comma_id: + frame["mode"] = "key_or_end" + self.mode = "object_key_or_end" + else: + self.stack.pop() + self._value_finished() + self.number_text = "" + return + self.number_text += self._regular_surfaces.get(token_id, "") + return + + if self.mode == "array_value_or_end": + if token_id == self.rbrack_id: + self.stack.pop() + self._value_finished() + return + self.mode = "expect_value" + self.type_hint = "any" + self.step(token_id) + return + + if self.mode == "array_after_value": + frame = self.stack[-1] + if token_id == self.comma_id: + frame["mode"] = "value" + self.mode = "expect_value" + self.type_hint = "any" + return + self.stack.pop() + self._value_finished() + return + + if self.mode == "object_key_or_end": + if token_id == self.rbrace_id: + self.stack.pop() + self._value_finished() + return + self.mode = "object_key_string" + return + + if self.mode == "object_key_string": + if token_id == self.quote_id: + self.mode = "object_after_key" + return + + if self.mode == "object_after_key": + self.mode = "expect_value" + self.type_hint = "any" + return + + if self.mode == "object_after_value": + frame = self.stack[-1] + if token_id == self.comma_id: + frame["mode"] = "key_or_end" + self.mode = "object_key_or_end" + return + self.stack.pop() + self._value_finished() + + def _is_value_complete(self): + return _is_complete_number(self.number_text) + + +class ToolCallCFG: + def __init__(self, tokenizer, tools_text): + self.tokenizer = tokenizer + self.tools_text = tools_text + self.tools = _normalize_tool_schema_spec(tools_text) + self.tool_names = [tool["name"] for tool in self.tools] + self.param_names = { + tool["name"]: [name for name, _ in tool["parameters"]] + for tool in self.tools + } + self.param_types = { + tool["name"]: {name: _schema_type_name(param_type) for name, param_type in tool["parameters"]} + for tool in self.tools + } + self.ids = tokenizer.json_token_ids + self.quote_id = self.ids[JSON_QUOTE] + self.tool_name_trie = _build_trie( + (i, tokenizer.encode_json_string_content(name)) + for i, name in enumerate(self.tool_names) + ) + self.state = "start_array" + self.current_tool = None + self.current_tool_index = None + self.current_param_start = 0 + self.current_param_name = None + self.current_param_type = "any" + self._name_trie_node = None + self._selected_name_index = None + self.value_parser = None + self.invalid = False + + def _remaining_param_trie(self): + if self.current_tool is None: + return {} + params = self.param_names.get(self.current_tool, []) + sequences = [] + for i in range(self.current_param_start, len(params)): + sequences.append((i, self.tokenizer.encode_json_string_content(params[i]))) + return _build_trie(sequences) + + def _set_name_state(self, trie, next_state): + self.state = next_state + self._name_trie_node = trie + self._selected_name_index = None + + def allowed_next_ids(self, training=False): + if self.invalid: + return None + if self.value_parser is not None: + return self.value_parser.allowed_next_ids(training=training) + + if self.state == "done": + return np.array([EOS_ID], dtype=np.int32) + if self.state == "start_array": + return np.array([self.ids[JSON_LBRACK]], dtype=np.int32) + if self.state == "after_array_start": + return np.array([self.ids[JSON_RBRACK], self.ids[JSON_LBRACE]], dtype=np.int32) + if self.state == "after_call": + return np.array([self.ids[JSON_COMMA], self.ids[JSON_RBRACK]], dtype=np.int32) + if self.state == "expect_key_name": + return np.array([self.ids[JSON_KEY_NAME]], dtype=np.int32) + if self.state == "expect_name_colon": + return np.array([self.ids[JSON_COLON]], dtype=np.int32) + if self.state == "expect_tool_name_quote": + return np.array([self.quote_id], dtype=np.int32) + if self.state == "tool_name_content": + allowed = [token_id for token_id in self._name_trie_node.keys() if token_id != "_end"] + if "_end" in self._name_trie_node: + allowed.append(self.quote_id) + return np.array(sorted(set(allowed)), dtype=np.int32) + if self.state == "expect_tool_name_comma": + return np.array([self.ids[JSON_COMMA]], dtype=np.int32) + if self.state == "expect_key_arguments": + return np.array([self.ids[JSON_KEY_ARGUMENTS]], dtype=np.int32) + if self.state == "expect_arguments_colon": + return np.array([self.ids[JSON_COLON]], dtype=np.int32) + if self.state == "expect_arguments_object": + return np.array([self.ids[JSON_LBRACE]], dtype=np.int32) + if self.state == "arg_name_or_end": + params = self.param_names.get(self.current_tool, []) + if self.current_param_start >= len(params): + return np.array([self.ids[JSON_RBRACE]], dtype=np.int32) + return np.array([self.ids[JSON_RBRACE], self.quote_id], dtype=np.int32) + if self.state == "arg_name_content": + allowed = [token_id for token_id in self._name_trie_node.keys() if token_id != "_end"] + if "_end" in self._name_trie_node: + allowed.append(self.quote_id) + return np.array(sorted(set(allowed)), dtype=np.int32) + if self.state == "expect_arg_colon": + return np.array([self.ids[JSON_COLON]], dtype=np.int32) + if self.state == "after_arg_value": + has_more = self.current_param_start < len(self.param_names.get(self.current_tool, [])) + if has_more: + return np.array([self.ids[JSON_COMMA], self.ids[JSON_RBRACE]], dtype=np.int32) + return np.array([self.ids[JSON_RBRACE]], dtype=np.int32) + if self.state == "expect_call_end": + return np.array([self.ids[JSON_RBRACE]], dtype=np.int32) + return None + + def step(self, token_id): + token_id = int(token_id) + if self.invalid: + return + if self.value_parser is not None: + self.value_parser.step(token_id) + if self.value_parser.complete: + self.value_parser = None + self.state = "after_arg_value" + return + + if self.state == "start_array": + self.state = "after_array_start" + return + if self.state == "after_array_start": + if token_id == self.ids[JSON_RBRACK]: + self.state = "done" + else: + self.state = "expect_key_name" + return + if self.state == "after_call": + if token_id == self.ids[JSON_COMMA]: + self.state = "expect_key_name" + else: + self.state = "done" + return + if self.state == "expect_key_name": + self.state = "expect_name_colon" + return + if self.state == "expect_name_colon": + self.state = "expect_tool_name_quote" + return + if self.state == "expect_tool_name_quote": + self._set_name_state(self.tool_name_trie, "tool_name_content") + return + if self.state == "tool_name_content": + if token_id == self.quote_id: + if "_end" not in self._name_trie_node: + self.invalid = True + return + self.current_tool_index = self._name_trie_node["_end"] + self.current_tool = self.tool_names[self.current_tool_index] + self.current_param_start = 0 + self.state = "expect_tool_name_comma" + else: + if token_id not in self._name_trie_node: + self.invalid = True + return + self._name_trie_node = self._name_trie_node[token_id] + return + if self.state == "expect_tool_name_comma": + self.state = "expect_key_arguments" + return + if self.state == "expect_key_arguments": + self.state = "expect_arguments_colon" + return + if self.state == "expect_arguments_colon": + self.state = "expect_arguments_object" + return + if self.state == "expect_arguments_object": + self.state = "arg_name_or_end" + return + if self.state == "arg_name_or_end": + if token_id == self.ids[JSON_RBRACE]: + self.state = "expect_call_end" + else: + self._set_name_state(self._remaining_param_trie(), "arg_name_content") + return + if self.state == "arg_name_content": + if token_id == self.quote_id: + if "_end" not in self._name_trie_node: + self.invalid = True + return + param_idx = self._name_trie_node["_end"] + self.current_param_name = self.param_names[self.current_tool][param_idx] + self.current_param_type = self.param_types[self.current_tool].get(self.current_param_name, "any") + self.current_param_start = param_idx + 1 + self.state = "expect_arg_colon" + else: + if token_id not in self._name_trie_node: + self.invalid = True + return + self._name_trie_node = self._name_trie_node[token_id] + return + if self.state == "expect_arg_colon": + self.value_parser = JsonValueParser(self.tokenizer, type_hint=self.current_param_type) + return + if self.state == "after_arg_value": + if token_id == self.ids[JSON_COMMA]: + self.state = "arg_name_or_end" + else: + self.state = "expect_call_end" + return + if self.state == "expect_call_end": + self.state = "after_call" + + +def _extract_tools_tokens(tgt_in, loss_mask): + supervised = np.flatnonzero(loss_mask > 0) + if len(supervised) == 0: + return [] + tool_stop = int(supervised[0]) + 1 + return [int(token_id) for token_id in tgt_in[2:tool_stop] if int(token_id) != 0] + + +def build_cfg_training_constraints(tokenizer, tgt_in_batch, tgt_out_batch, loss_mask_batch): + batch_size, seq_len = tgt_out_batch.shape + batch_allowed = [[None] * seq_len for _ in range(batch_size)] + max_allowed = 0 + + for i in range(batch_size): + tools_tokens = _extract_tools_tokens(tgt_in_batch[i], loss_mask_batch[i]) + tools_text = tokenizer.decode_structured(tools_tokens) + cfg = ToolCallCFG(tokenizer, tools_text) + active_positions = np.flatnonzero(loss_mask_batch[i] > 0) + for pos in active_positions: + if cfg.invalid: + break + allowed = cfg.allowed_next_ids(training=True) + if allowed is not None and len(allowed) > 0: + allowed = np.asarray(allowed, dtype=np.int32) + batch_allowed[i][int(pos)] = allowed + max_allowed = max(max_allowed, len(allowed)) + cfg.step(int(tgt_out_batch[i, pos])) + + if max_allowed == 0: + return ( + np.full((batch_size, seq_len, 1), -1, dtype=np.int32), + np.zeros((batch_size, seq_len), dtype=np.int32), + ) + + allowed_ids = np.full((batch_size, seq_len, max_allowed), -1, dtype=np.int32) + allowed_counts = np.zeros((batch_size, seq_len), dtype=np.int32) + for i in range(batch_size): + for pos, allowed in enumerate(batch_allowed[i]): + if allowed is None: + continue + count = len(allowed) + allowed_ids[i, pos, :count] = allowed + allowed_counts[i, pos] = count + return allowed_ids, allowed_counts diff --git a/src/toucan.py b/src/toucan.py new file mode 100644 index 0000000..90bee46 --- /dev/null +++ b/src/toucan.py @@ -0,0 +1,149 @@ +import hashlib +import json +import os + +from datasets import load_dataset +import numpy as np + +from .data import CACHE_DIR + +TOUCAN_DATASET = "Agent-Ark/Toucan-1.5M" + + +def parse_target_tools(text): + return [name.strip() for name in text.split(",") if name.strip()] + + +def _tool_type_text(spec): + if isinstance(spec, dict) and "type" in spec and isinstance(spec["type"], str): + return spec["type"] + return "any" + + +def _tool_desc_text(spec): + if isinstance(spec, dict) and "description" in spec and isinstance(spec["description"], str): + return spec["description"].strip() + return "" + + +def _tool_properties(tool): + params = tool["function"]["parameters"] + if isinstance(params, dict) and "properties" in params and isinstance(params["properties"], dict): + return params["properties"] + return {} + + +def _tool_required(tool): + params = tool["function"]["parameters"] + if isinstance(params, dict) and "required" in params and isinstance(params["required"], list): + return set(params["required"]) + return set() + + +def _tool_matches_target(tool_name, target_name): + return ( + tool_name == target_name + or tool_name.endswith(f"-{target_name}") + or tool_name.endswith(f"::{target_name}") + ) + + +def compact_toucan_tool(tool): + properties = _tool_properties(tool) + return { + "name": tool["function"]["name"], + "parameters": {name: _tool_type_text(spec) for name, spec in properties.items()}, + } + + +def format_toucan_tool(tool): + function = tool["function"] + properties = _tool_properties(tool) + required = _tool_required(tool) + lines = [ + f"name: {function['name']}", + f"description: {_tool_desc_text(function)}", + "parameters:", + ] + for name, spec in properties.items(): + req = " required" if name in required else "" + desc = _tool_desc_text(spec) + line = f"- {name}: {_tool_type_text(spec)}{req}" + if desc: + line = f"{line} | {desc}" + lines.append(line) + return "\n".join(lines) + + +def prepare_toucan_example(row): + tools = json.loads(row["available_tools"]) + compact_tools = [compact_toucan_tool(tool) for tool in tools] + tool_texts = [format_toucan_tool(tool) for tool in tools] + target_tools = parse_target_tools(row["target_tools"]) + positive_indices = [ + i for i, tool in enumerate(compact_tools) + if any(_tool_matches_target(tool["name"], target_name) for target_name in target_tools) + ] + return { + "subset_name": row["subset_name"], + "question": row["question"], + "target_tools": target_tools, + "tool_names": [tool["name"] for tool in compact_tools], + "positive_indices": positive_indices, + "tools_json": json.dumps(compact_tools, ensure_ascii=True, separators=(",", ":")), + "tool_texts": tool_texts, + } + + +def _toucan_cache_path(config, split, max_samples): + cache_id = hashlib.md5(f"toucan_{config}_{split}_{max_samples}".encode()).hexdigest()[:12] + return os.path.join(CACHE_DIR, f"{cache_id}_toucan.jsonl") + + +def cache_toucan_examples(config="Kimi-K2", split="train", max_samples=None, tokenizer=None, max_text_len=256): + os.makedirs(CACHE_DIR, exist_ok=True) + path = _toucan_cache_path(config, split, max_samples) + if os.path.exists(path): + return path + + ds = load_dataset(TOUCAN_DATASET, config, split=split, streaming=True) + with open(path, "w") as f: + for i, row in enumerate(ds): + if max_samples is not None and i >= max_samples: + break + ex = prepare_toucan_example(row) + if tokenizer is not None: + ex["question_ids"] = tokenizer.encode(ex["question"])[:max_text_len] + ex["tool_ids"] = [tokenizer.encode(text)[:max_text_len] for text in ex["tool_texts"]] + f.write(json.dumps(ex, ensure_ascii=True) + "\n") + return path + + +def load_toucan_contrastive_data(path, max_text_len=256): + rows = [json.loads(line) for line in open(path)] + rows = [row for row in rows if row["positive_indices"] and len(row["positive_indices"]) < len(row["tool_names"])] + max_tools = max(len(row["tool_ids"]) for row in rows) + questions = np.zeros((len(rows), max_text_len), dtype=np.int32) + tools = np.zeros((len(rows), max_tools, max_text_len), dtype=np.int32) + labels = np.zeros((len(rows), max_tools), dtype=np.float32) + tool_mask = np.zeros((len(rows), max_tools), dtype=np.float32) + + for i, row in enumerate(rows): + q_ids = row["question_ids"][:max_text_len] + questions[i, :len(q_ids)] = q_ids + positive = set(row["positive_indices"]) + for j, tool_ids in enumerate(row["tool_ids"][:max_tools]): + tool_ids = tool_ids[:max_text_len] + tools[i, j, :len(tool_ids)] = tool_ids + labels[i, j] = 1.0 if j in positive else 0.0 + tool_mask[i, j] = 1.0 + return questions, tools, labels, tool_mask + + +def get_toucan_batches(questions, tools, labels, tool_mask, batch_size, shuffle=True): + n = len(questions) + while True: + indices = np.random.permutation(n) if shuffle else np.arange(n) + for i in range(0, n - batch_size + 1, batch_size): + idx = indices[i:i + batch_size] + yield questions[idx], tools[idx], labels[idx], tool_mask[idx] diff --git a/src/train.py b/src/train.py index 92ddf57..37a440d 100644 --- a/src/train.py +++ b/src/train.py @@ -14,7 +14,7 @@ from flax.training import train_state from .data import ( - get_batches, get_tokenizer, get_speech_batches, + get_batches, get_tokenizer, get_speech_batches, get_text_mel_batches, load_prepared_data, load_prepared_mels, load_example_with_audio, PrefetchIterator, count_batches, @@ -26,6 +26,7 @@ make_padding_mask, make_mel_padding_mask, ) +from .toucan import get_toucan_batches, load_toucan_contrastive_data def _newton_schulz(G, steps=5): """Approximate polar decomposition via Newton-Schulz iteration.""" @@ -227,9 +228,59 @@ def _maybe_quantize(path, leaf): _MAT_FACTORS = () _MAT_FF_WIDTHS = () _D_FF = 2048 +_TOOL_CONTRASTIVE_WEIGHT = 0.0 +_AUDIO_TEXT_CONTRASTIVE_WEIGHT = 0.0 -def _text_loss_fn(state, params, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask): +def _tool_contrastive_loss(state, params, contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask): + pad_id = 0 + q_mask = make_padding_mask(contrastive_q, pad_id) + q_slots = state.apply_fn({"params": _quantize_params(params, group_size=_GROUP_SIZE)}, + contrastive_q, src_mask=q_mask, method="encode").astype(jnp.float32) + + bsz, n_tools, seq_len = contrastive_tools.shape + flat_tools = contrastive_tools.reshape(bsz * n_tools, seq_len) + tool_mask = make_padding_mask(flat_tools, pad_id) + t_slots = state.apply_fn({"params": _quantize_params(params, group_size=_GROUP_SIZE)}, + flat_tools, src_mask=tool_mask, method="encode").astype(jnp.float32) + t_slots = t_slots.reshape(bsz, n_tools, q_slots.shape[1], q_slots.shape[2]) + + scores = jnp.sum(q_slots[:, None, :, :] * t_slots, axis=(2, 3)) / q_slots.shape[1] + loss = optax.sigmoid_binary_cross_entropy(scores, contrastive_labels) + return jnp.sum(loss * contrastive_tool_mask) / jnp.maximum(jnp.sum(contrastive_tool_mask), 1.0) + + +def _siglip_pair_loss(scores): + labels = 2.0 * jnp.eye(scores.shape[0], dtype=scores.dtype) - 1.0 + return jnp.mean(jax.nn.softplus(-labels * scores)) + + +def _audio_text_contrastive_loss(state, params, contrastive_text, contrastive_mel): + pad_id = 0 + text_mask = make_padding_mask(contrastive_text, pad_id) + text_slots = state.apply_fn( + {"params": _quantize_params(params, group_size=_GROUP_SIZE)}, + contrastive_text, + src_mask=text_mask, + method="encode", + ).astype(jnp.float32) + + mel_mask = make_mel_padding_mask(contrastive_mel) + audio_slots = state.apply_fn( + {"params": _quantize_params(params, group_size=_GROUP_SIZE)}, + contrastive_mel, + src_mask=mel_mask, + deterministic=True, + method="encode_speech", + ).astype(jnp.float32) + + scores = jnp.einsum("bmd,cmd->bc", text_slots, audio_slots) / text_slots.shape[1] + return _siglip_pair_loss(scores) + + +def _text_loss_fn(state, params, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask, + contrastive_q=None, contrastive_tools=None, contrastive_labels=None, contrastive_tool_mask=None, + audio_text=None, audio_mel=None): pad_id = 0 src_mask = make_padding_mask(src, pad_id) tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) @@ -243,12 +294,19 @@ def _text_loss_fn(state, params, src, tgt_in, tgt_out, causal_mask, ffn_mask, rn ) logits_f32 = logits.astype(jnp.float32) mask = loss_mask - ce_loss = jnp.sum( - optax.softmax_cross_entropy_with_integer_labels(logits_f32, tgt_out) * mask - ) / jnp.maximum(jnp.sum(mask), 1.0) + token_loss = optax.softmax_cross_entropy_with_integer_labels(logits_f32, tgt_out) + ce_loss = jnp.sum(token_loss * mask) / jnp.maximum(jnp.sum(mask), 1.0) z_loss = 1e-4 * jnp.mean(jax.nn.logsumexp(logits_f32, axis=-1) ** 2) div_loss = 1e-4 * slot_div - return ce_loss + z_loss + div_loss + tc_loss = 0.0 + if _TOOL_CONTRASTIVE_WEIGHT > 0 and contrastive_q is not None: + tc_loss = _tool_contrastive_loss( + state, params, contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask + ) + at_loss = 0.0 + if _AUDIO_TEXT_CONTRASTIVE_WEIGHT > 0 and audio_text is not None: + at_loss = _audio_text_contrastive_loss(state, params, audio_text, audio_mel) + return ce_loss + z_loss + div_loss + _TOOL_CONTRASTIVE_WEIGHT * tc_loss + _AUDIO_TEXT_CONTRASTIVE_WEIGHT * at_loss def _speech_loss_fn(state, params, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask): @@ -266,9 +324,8 @@ def _speech_loss_fn(state, params, mel, tgt_in, tgt_out, causal_mask, ffn_mask, ) logits_f32 = logits.astype(jnp.float32) mask = loss_mask - ce_loss = jnp.sum( - optax.softmax_cross_entropy_with_integer_labels(logits_f32, tgt_out) * mask - ) / jnp.maximum(jnp.sum(mask), 1.0) + token_loss = optax.softmax_cross_entropy_with_integer_labels(logits_f32, tgt_out) + ce_loss = jnp.sum(token_loss * mask) / jnp.maximum(jnp.sum(mask), 1.0) z_loss = 1e-4 * jnp.mean(jax.nn.logsumexp(logits_f32, axis=-1) ** 2) div_loss = 1e-4 * slot_div return ce_loss + z_loss + div_loss @@ -292,10 +349,16 @@ def _make_ffn_mask(batch_size, d_ff, mat_ff_widths): return jnp.concatenate(rows, axis=0) -def _train_step_text(state, ema_params, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask): +def _train_step_text(state, ema_params, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask, + contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask, + audio_text, audio_mel): ema_decay = 0.999 loss, grads = jax.value_and_grad( - lambda p: _text_loss_fn(state, p, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask) + lambda p: _text_loss_fn( + state, p, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask, + contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask, + audio_text, audio_mel, + ) )(state.params) grads = jax.lax.pmean(grads, axis_name="batch") loss = jax.lax.pmean(loss, axis_name="batch") @@ -305,11 +368,18 @@ def _train_step_text(state, ema_params, src, tgt_in, tgt_out, causal_mask, ffn_m return state, ema_params, loss, grad_norm -def _train_step_text_masked(state, ema_params, src, tgt_in, tgt_out, causal_mask, prune_mask, ffn_mask, rng, loss_mask): +def _train_step_text_masked(state, ema_params, src, tgt_in, tgt_out, causal_mask, prune_mask, + ffn_mask, rng, loss_mask, + contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask, + audio_text, audio_mel): """Text training step with fused prune mask application.""" ema_decay = 0.999 loss, grads = jax.value_and_grad( - lambda p: _text_loss_fn(state, p, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask) + lambda p: _text_loss_fn( + state, p, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask, + contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask, + audio_text, audio_mel, + ) )(state.params) grads = jax.lax.pmean(grads, axis_name="batch") loss = jax.lax.pmean(loss, axis_name="batch") @@ -324,7 +394,9 @@ def _train_step_text_masked(state, ema_params, src, tgt_in, tgt_out, causal_mask def _train_step_speech(state, ema_params, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask): ema_decay = 0.999 loss, grads = jax.value_and_grad( - lambda p: _speech_loss_fn(state, p, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask) + lambda p: _speech_loss_fn( + state, p, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask, + ) )(state.params) grads = jax.lax.pmean(grads, axis_name="batch") loss = jax.lax.pmean(loss, axis_name="batch") @@ -334,11 +406,14 @@ def _train_step_speech(state, ema_params, mel, tgt_in, tgt_out, causal_mask, ffn return state, ema_params, loss, grad_norm -def _train_step_speech_masked(state, ema_params, mel, tgt_in, tgt_out, causal_mask, prune_mask, ffn_mask, rng, loss_mask): +def _train_step_speech_masked(state, ema_params, mel, tgt_in, tgt_out, causal_mask, prune_mask, + ffn_mask, rng, loss_mask): """Speech training step with fused prune mask application.""" ema_decay = 0.999 loss, grads = jax.value_and_grad( - lambda p: _speech_loss_fn(state, p, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask) + lambda p: _speech_loss_fn( + state, p, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask, + ) )(state.params) grads = jax.lax.pmean(grads, axis_name="batch") loss = jax.lax.pmean(loss, axis_name="batch") @@ -380,7 +455,6 @@ def val_loss_batch(params, src, tgt_in, tgt_out, causal_mask, loss_mask): return jnp.sum(loss * loss_mask), jnp.sum(loss_mask) return val_loss_batch - def _make_mat_val_loss_fn(apply_fn, ff_width): """Val loss for matryoshka sub-model at given FFN width.""" @jax.jit @@ -417,7 +491,6 @@ def val_loss_batch(params, mel, tgt_in, tgt_out, causal_mask, loss_mask): return jnp.sum(loss * mask), jnp.sum(mask) return val_loss_batch - def _estimate_mat_params(config, matryoshka_factor): """Estimate parameter count of a sub-model at a given matryoshka factor. @@ -520,9 +593,11 @@ def train(args): dropout_rate=getattr(args, "dropout", 0.1), ) - global _GROUP_SIZE, _MAT_FACTORS, _MAT_FF_WIDTHS, _D_FF + global _GROUP_SIZE, _MAT_FACTORS, _MAT_FF_WIDTHS, _D_FF, _TOOL_CONTRASTIVE_WEIGHT, _AUDIO_TEXT_CONTRASTIVE_WEIGHT _GROUP_SIZE = getattr(args, "group_size", 32) _D_FF = config.d_ff + _TOOL_CONTRASTIVE_WEIGHT = getattr(args, "tool_contrastive_weight", 0.0) + _AUDIO_TEXT_CONTRASTIVE_WEIGHT = getattr(args, "audio_text_contrastive_weight", 0.0) mat_factors_raw = getattr(args, "mat_factors", None) if mat_factors_raw: _MAT_FACTORS = tuple(f for f in mat_factors_raw if f > 1) @@ -578,6 +653,9 @@ def train(args): print(f" Activation {config.activation:>12}") print(f" Dtype {config.dtype:>12}") print(f" Dropout {config.dropout_rate:>12}") + print(f" CFG infer {str(getattr(args, 'cfg_inference', False)):>12}") + print(f" Tool ctr wt {_TOOL_CONTRASTIVE_WEIGHT:>12.4f}") + print(f" Audio/text wt {_AUDIO_TEXT_CONTRASTIVE_WEIGHT:>12.4f}") if not no_speech and train_mels is not None: print(f" Speech {speech_batches_per_epoch} batches/epoch") print(f" n_mels {n_mels:>12}") @@ -627,6 +705,12 @@ def train(args): prune_interval = getattr(args, "prune_interval", 100) prune_start_frac = getattr(args, "prune_start_frac", 0.33) prune_end_frac = getattr(args, "prune_end_frac", 0.67) + toucan_batch = None + if _TOOL_CONTRASTIVE_WEIGHT > 0: + toucan_path = train_data.get("toucan_cache_path") + if not toucan_path: + raise FileNotFoundError("No Toucan cache found. Re-run `needle tokenize --toucan-config ...` first.") + toucan_batch = load_toucan_contrastive_data(toucan_path, max_text_len=args.max_enc_len) weight_prune_epoch = 0 if sparsity_ratio > 0 else -1 @@ -645,6 +729,19 @@ def train(args): loss_mask=train_loss_mask), prefetch=4, ) + contrastive_iter = None + if toucan_batch is not None: + contrastive_iter = PrefetchIterator( + lambda: get_toucan_batches(*toucan_batch, unique_batch_size), + prefetch=2, + ) + + audio_text_iter = None + if _AUDIO_TEXT_CONTRASTIVE_WEIGHT > 0 and train_mels is not None: + audio_text_iter = PrefetchIterator( + lambda: get_text_mel_batches(enc_inputs, train_mels, unique_batch_size), + prefetch=2, + ) speech_batch_iter = None if not no_speech and train_mels is not None: @@ -678,6 +775,18 @@ def train(args): if do_text: src, tgt_in, tgt_out, lm = next(text_batch_iter) text_idx += 1 + if contrastive_iter is not None: + tc_q, tc_tools, tc_labels, tc_tool_mask = next(contrastive_iter) + else: + tc_q = np.zeros((len(tgt_in), args.max_enc_len), dtype=np.int32) + tc_tools = np.zeros((len(tgt_in), 1, args.max_enc_len), dtype=np.int32) + tc_labels = np.zeros((len(tgt_in), 1), dtype=np.float32) + tc_tool_mask = np.zeros((len(tgt_in), 1), dtype=np.float32) + if audio_text_iter is not None: + at_text, at_mel = next(audio_text_iter) + else: + at_text = np.zeros((len(tgt_in), args.max_enc_len), dtype=np.int32) + at_mel = np.zeros((len(tgt_in), args.max_mel_len, args.n_mels), dtype=np.float32) if n_widths > 1 and mat_shared_input: per_width = args.batch_size // n_widths @@ -688,22 +797,40 @@ def _tile_for_mat(arr): tgt_in_b = _tile_for_mat(tgt_in) tgt_out_b = _tile_for_mat(tgt_out) lm_b = _tile_for_mat(lm) + tc_q_b = _tile_for_mat(tc_q) + tc_tools_b = _tile_for_mat(tc_tools) + tc_labels_b = _tile_for_mat(tc_labels) + tc_tool_mask_b = _tile_for_mat(tc_tool_mask) + at_text_b = _tile_for_mat(at_text) + at_mel_b = _tile_for_mat(at_mel) else: src_b = shard_batch(src, num_devices) tgt_in_b = shard_batch(tgt_in, num_devices) tgt_out_b = shard_batch(tgt_out, num_devices) lm_b = shard_batch(lm, num_devices) + tc_q_b = shard_batch(tc_q, num_devices) + tc_tools_b = shard_batch(tc_tools, num_devices) + tc_labels_b = shard_batch(tc_labels, num_devices) + tc_tool_mask_b = shard_batch(tc_tool_mask, num_devices) + at_text_b = shard_batch(at_text, num_devices) + at_mel_b = shard_batch(at_mel, num_devices) rng, text_rng = jax.random.split(rng) text_rngs = jax.random.split(text_rng, num_devices) if prune_mask is not None: state, ema_params, loss, grad_norm = p_train_step_masked( - state, ema_params, src_b, tgt_in_b, tgt_out_b, causal_mask, prune_mask, text_ffn_mask, text_rngs, lm_b, + state, ema_params, src_b, tgt_in_b, tgt_out_b, causal_mask, prune_mask, + text_ffn_mask, text_rngs, lm_b, + tc_q_b, tc_tools_b, tc_labels_b, tc_tool_mask_b, + at_text_b, at_mel_b, ) else: state, ema_params, loss, grad_norm = p_train_step( - state, ema_params, src_b, tgt_in_b, tgt_out_b, causal_mask, text_ffn_mask, text_rngs, lm_b, + state, ema_params, src_b, tgt_in_b, tgt_out_b, causal_mask, + text_ffn_mask, text_rngs, lm_b, + tc_q_b, tc_tools_b, tc_labels_b, tc_tool_mask_b, + at_text_b, at_mel_b, ) text_loss_val = float(loss[0]) @@ -735,11 +862,13 @@ def _tile_sp(arr): if prune_mask is not None: state, ema_params, sp_loss, sp_grad_norm = p_train_step_speech_masked( - state, ema_params, mel_b, sp_tgt_in_b, sp_tgt_out_b, causal_mask, prune_mask, text_ffn_mask, spec_rngs, sp_lm_b, + state, ema_params, mel_b, sp_tgt_in_b, sp_tgt_out_b, causal_mask, prune_mask, + text_ffn_mask, spec_rngs, sp_lm_b, ) else: state, ema_params, sp_loss, sp_grad_norm = p_train_step_speech( - state, ema_params, mel_b, sp_tgt_in_b, sp_tgt_out_b, causal_mask, text_ffn_mask, spec_rngs, sp_lm_b, + state, ema_params, mel_b, sp_tgt_in_b, sp_tgt_out_b, causal_mask, + text_ffn_mask, spec_rngs, sp_lm_b, ) speech_loss_val = float(sp_loss[0]) speech_losses.append(speech_loss_val) @@ -790,8 +919,8 @@ def _tile_sp(arr): "train/adam_lr": float(adam_schedule(global_step)), "train/muon_lr": float(muon_schedule(global_step)), "train/tokens_per_sec": tokens_per_batch / dt, - "train/step": global_step, - } + "train/step": global_step, + } if speech_loss_val is not None: log_dict["train/speech_loss"] = speech_loss_val if epoch == weight_prune_epoch and not gradual_sparsify_done: @@ -801,6 +930,10 @@ def _tile_sp(arr): wandb.log(log_dict) text_batch_iter.close() + if contrastive_iter is not None: + contrastive_iter.close() + if audio_text_iter is not None: + audio_text_iter.close() if speech_batch_iter is not None: speech_batch_iter.close() @@ -826,6 +959,8 @@ def _tile_sp(arr): final_loss = text_losses[-1] if text_losses else float("nan") final_ppl = math.exp(min(final_loss, 20)) if not math.isnan(final_loss) else float("nan") + epoch_eval_t0 = time.perf_counter() + print(f"\n [epoch {epoch + 1}] starting epoch-end evaluation...") eval_params = jax_utils.unreplicate(ema_params) val_causal = make_causal_mask(args.max_dec_len) @@ -841,8 +976,7 @@ def _tile_sp(arr): q_loss, q_toks = 0.0, 0.0 mat_accum = {f: [0.0, 0.0] for f in _MAT_FACTORS} - for vb in get_batches(val_enc, val_dec_in, val_dec_tgt, args.batch_size, - shuffle=False, loss_mask=val_loss_mask): + for vb in get_batches(val_enc, val_dec_in, val_dec_tgt, args.batch_size, shuffle=False, loss_mask=val_loss_mask): src, dec_in, dec_tgt, lm = vb[0], vb[1], vb[2], vb[3] vl, vt = val_loss_fn(eval_params, src, dec_in, dec_tgt, val_causal, lm) full_loss += float(vl); full_toks += float(vt) @@ -856,6 +990,7 @@ def _tile_sp(arr): last_val_ppl = float(math.exp(min(full_loss / max(full_toks, 1), 20))) quant_val_ppl = float(math.exp(min(q_loss / max(q_toks, 1), 20))) del q_params + print(f" [epoch {epoch + 1}] text/quant/mat val done in {time.perf_counter() - epoch_eval_t0:.1f}s") mat_results = {} for f in _MAT_FACTORS: @@ -866,13 +1001,13 @@ def _tile_sp(arr): speech_val_ppl = None if speech_vl_fn is not None and val_mels is not None: sp_total_loss, sp_total_toks = 0.0, 0.0 - for sp_batch in get_speech_batches(val_mels, val_dec_in, val_dec_tgt, args.batch_size, - shuffle=False, loss_mask=val_loss_mask): + for sp_batch in get_speech_batches(val_mels, val_dec_in, val_dec_tgt, args.batch_size, shuffle=False, loss_mask=val_loss_mask): vl, vt = speech_vl_fn(eval_params, sp_batch[0], sp_batch[1], sp_batch[2], val_causal, sp_batch[3]) sp_total_loss += float(vl) sp_total_toks += float(vt) speech_val_loss = sp_total_loss / max(sp_total_toks, 1) speech_val_ppl = float(math.exp(min(speech_val_loss, 20))) + print(f" [epoch {epoch + 1}] speech val done in {time.perf_counter() - epoch_eval_t0:.1f}s") params_np = jax.tree.map(np.array, eval_params) total_params = sum(x.size for x in jax.tree.leaves(params_np)) @@ -884,10 +1019,12 @@ def _tile_sp(arr): with open(ckpt_path, "wb") as f: pickle.dump({"params": params_np, "config": config.__dict__}, f) del params_np + print(f" [epoch {epoch + 1}] checkpoint saved in {time.perf_counter() - epoch_eval_t0:.1f}s") from .test import measure_throughput from .run import generate, generate_from_audio tp = measure_throughput(eval_model, eval_params, tokenizer, num_runs=5) + print(f" [epoch {epoch + 1}] throughput benchmark done in {time.perf_counter() - epoch_eval_t0:.1f}s") from .data import load_tool_calls _, val_global_indices = load_tool_calls( @@ -913,6 +1050,7 @@ def _tile_sp(arr): text_pred = generate( eval_model, eval_params, tokenizer, pair["query"], tools=pair["tools"], max_gen_len=args.max_dec_len, seed=i, stream=False, + use_cfg=getattr(args, "cfg_inference", False), ).strip()[:120] # Voice prediction @@ -921,9 +1059,11 @@ def _tile_sp(arr): voice_pred = generate_from_audio( eval_model, eval_params, tokenizer, pair["audio_array"], sr=pair["sampling_rate"], tools=pair["tools"], max_gen_len=args.max_dec_len, seed=i, stream=False, + use_cfg=getattr(args, "cfg_inference", False), ).strip()[:120] unified_samples.append((query, ref, text_pred, voice_pred)) + print(f" [epoch {epoch + 1}] sample generation done in {time.perf_counter() - epoch_eval_t0:.1f}s") del eval_params @@ -983,4 +1123,3 @@ def _tile_sp(arr): if use_wandb: wandb.finish() print("\nTraining complete.") - From 007178a0f83edd6fd499c85516cc3e73840838c3 Mon Sep 17 00:00:00 2001 From: Karen Mosoyan Date: Tue, 10 Mar 2026 07:02:56 +0000 Subject: [PATCH 3/6] clean --- scripts/inspect_toucan.py | 36 ------------------- src/train.py | 75 +++++++++++++++++++-------------------- 2 files changed, 37 insertions(+), 74 deletions(-) delete mode 100644 scripts/inspect_toucan.py diff --git a/scripts/inspect_toucan.py b/scripts/inspect_toucan.py deleted file mode 100644 index 96a4a65..0000000 --- a/scripts/inspect_toucan.py +++ /dev/null @@ -1,36 +0,0 @@ -import argparse -import json - -from datasets import load_dataset - -from src.toucan import prepare_toucan_example - - -def parse_args(): - parser = argparse.ArgumentParser() - parser.add_argument("--config", type=str, default="Kimi-K2") - parser.add_argument("--split", type=str, default="train") - parser.add_argument("--samples", type=int, default=5) - return parser.parse_args() - - -def main(): - args = parse_args() - ds = load_dataset("Agent-Ark/Toucan-1.5M", args.config, split=args.split, streaming=True) - for i, row in enumerate(ds): - if i >= args.samples: - break - ex = prepare_toucan_example(row) - print(f"ROW {i}") - print(f"subset_name: {ex['subset_name']}") - print(f"question: {ex['question'][:160]}") - print(f"target_tools: {ex['target_tools']}") - print(f"positive_indices: {ex['positive_indices']}") - print(f"tool_names: {ex['tool_names'][:5]}") - print(f"tools_json: {ex['tools_json'][:200]}") - print(f"tool_text[0]: {ex['tool_texts'][0][:300]}") - print() - - -if __name__ == "__main__": - main() diff --git a/src/train.py b/src/train.py index 37a440d..3db8598 100644 --- a/src/train.py +++ b/src/train.py @@ -1021,49 +1021,48 @@ def _tile_sp(arr): del params_np print(f" [epoch {epoch + 1}] checkpoint saved in {time.perf_counter() - epoch_eval_t0:.1f}s") - from .test import measure_throughput - from .run import generate, generate_from_audio - tp = measure_throughput(eval_model, eval_params, tokenizer, num_runs=5) - print(f" [epoch {epoch + 1}] throughput benchmark done in {time.perf_counter() - epoch_eval_t0:.1f}s") - - from .data import load_tool_calls - _, val_global_indices = load_tool_calls( - "validation", - max_samples=val_data.get("split_max_samples"), - return_global_indices=True, - shuffle_before_split=val_data.get("shuffle_before_split", False), - shuffle_seed=val_data.get("split_seed", 42), - ) - val_kept = val_data["kept_indices"] - n_eval_samples = min(3, len(val_kept)) - step = max(1, len(val_kept) // n_eval_samples) - sample_indices = [val_kept[i * step] for i in range(n_eval_samples)] - eval_indices = val_global_indices[np.array(sample_indices)] - + tp = {"tokens_per_second": float("nan"), "avg_latency_s": float("nan")} unified_samples = [] - for i, ds_idx in enumerate(eval_indices): - pair = load_example_with_audio(int(ds_idx)) - query = pair["query"][:80] - ref = pair["answers"][:120] - - # Text prediction - text_pred = generate( - eval_model, eval_params, tokenizer, pair["query"], - tools=pair["tools"], max_gen_len=args.max_dec_len, seed=i, stream=False, - use_cfg=getattr(args, "cfg_inference", False), - ).strip()[:120] - - # Voice prediction - voice_pred = None - if not no_speech and pair["audio_array"] is not None: - voice_pred = generate_from_audio( - eval_model, eval_params, tokenizer, pair["audio_array"], sr=pair["sampling_rate"], + if not getattr(args, "skip_epoch_extras", False): + from .test import measure_throughput + from .run import generate, generate_from_audio + tp = measure_throughput(eval_model, eval_params, tokenizer, num_runs=5) + print(f" [epoch {epoch + 1}] throughput benchmark done in {time.perf_counter() - epoch_eval_t0:.1f}s") + + from .data import load_tool_calls + _, val_global_indices = load_tool_calls( + "validation", + max_samples=val_data.get("split_max_samples"), + return_global_indices=True, + shuffle_before_split=val_data.get("shuffle_before_split", False), + shuffle_seed=val_data.get("split_seed", 42), + ) + val_kept = val_data["kept_indices"] + n_eval_samples = min(3, len(val_kept)) + step = max(1, len(val_kept) // n_eval_samples) + sample_indices = [val_kept[i * step] for i in range(n_eval_samples)] + eval_indices = val_global_indices[np.array(sample_indices)] + + for i, ds_idx in enumerate(eval_indices): + pair = load_example_with_audio(int(ds_idx)) + query = pair["query"][:80] + ref = pair["answers"][:120] + text_pred = generate( + eval_model, eval_params, tokenizer, pair["query"], tools=pair["tools"], max_gen_len=args.max_dec_len, seed=i, stream=False, use_cfg=getattr(args, "cfg_inference", False), ).strip()[:120] - unified_samples.append((query, ref, text_pred, voice_pred)) - print(f" [epoch {epoch + 1}] sample generation done in {time.perf_counter() - epoch_eval_t0:.1f}s") + voice_pred = None + if not no_speech and pair["audio_array"] is not None: + voice_pred = generate_from_audio( + eval_model, eval_params, tokenizer, pair["audio_array"], sr=pair["sampling_rate"], + tools=pair["tools"], max_gen_len=args.max_dec_len, seed=i, stream=False, + use_cfg=getattr(args, "cfg_inference", False), + ).strip()[:120] + + unified_samples.append((query, ref, text_pred, voice_pred)) + print(f" [epoch {epoch + 1}] sample generation done in {time.perf_counter() - epoch_eval_t0:.1f}s") del eval_params From 0334bf2ebe4ab510182352677fcb04f1addd68ee Mon Sep 17 00:00:00 2001 From: Karen Mosoyan Date: Tue, 10 Mar 2026 23:40:42 +0000 Subject: [PATCH 4/6] added changes --- .../grid25_20260310_072111/summary.txt | 6 ++ .../grid25_20260310_072312/summary.txt | 32 ++++++ scripts/run_contrastive_sweep.py | 97 +++++++++++++++++++ scripts/summarize_sweep_logs.py | 53 ++++++++++ src/cli.py | 2 + src/train.py | 12 ++- 6 files changed, 197 insertions(+), 5 deletions(-) create mode 100644 logs/contrastive_sweeps/grid25_20260310_072111/summary.txt create mode 100644 logs/contrastive_sweeps/grid25_20260310_072312/summary.txt create mode 100644 scripts/run_contrastive_sweep.py create mode 100644 scripts/summarize_sweep_logs.py diff --git a/logs/contrastive_sweeps/grid25_20260310_072111/summary.txt b/logs/contrastive_sweeps/grid25_20260310_072111/summary.txt new file mode 100644 index 0000000..e1d821f --- /dev/null +++ b/logs/contrastive_sweeps/grid25_20260310_072111/summary.txt @@ -0,0 +1,6 @@ +mode=grid25 +epochs=1 +eval_every=1000000 +skip_epoch_extras=True +cfg_inference=False + diff --git a/logs/contrastive_sweeps/grid25_20260310_072312/summary.txt b/logs/contrastive_sweeps/grid25_20260310_072312/summary.txt new file mode 100644 index 0000000..ac0b234 --- /dev/null +++ b/logs/contrastive_sweeps/grid25_20260310_072312/summary.txt @@ -0,0 +1,32 @@ +mode=grid25 +epochs=1 +eval_every=1000000 +skip_epoch_extras=True +no_checkpoints=True +cfg_inference=False + +at_0p0__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_0p0.log +at_0p0__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_0p1.log +at_0p0__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_0p5.log +at_0p0__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_1p0.log +at_0p0__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_2p0.log +at_0p1__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_0p0.log +at_0p1__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_0p1.log +at_0p1__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_0p5.log +at_0p1__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_1p0.log +at_0p1__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_2p0.log +at_0p5__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_0p0.log +at_0p5__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_0p1.log +at_0p5__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_0p5.log +at_0p5__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_1p0.log +at_0p5__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_2p0.log +at_1p0__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_0p0.log +at_1p0__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_0p1.log +at_1p0__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_0p5.log +at_1p0__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_1p0.log +at_1p0__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_2p0.log +at_2p0__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_0p0.log +at_2p0__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_0p1.log +at_2p0__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_0p5.log +at_2p0__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_1p0.log +at_2p0__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_2p0.log diff --git a/scripts/run_contrastive_sweep.py b/scripts/run_contrastive_sweep.py new file mode 100644 index 0000000..c087ab8 --- /dev/null +++ b/scripts/run_contrastive_sweep.py @@ -0,0 +1,97 @@ +import argparse +import os +import subprocess +import sys +from datetime import datetime + + +WEIGHTS = [0.0, 0.1, 0.5, 1.0, 2.0] + + +def _fmt_weight(value): + text = str(value) + return text.replace(".", "p") + + +def _needle_path(repo_root): + local = os.path.join(repo_root, ".venv", "bin", "needle") + return local if os.path.exists(local) else "needle" + + +def _run_one(repo_root, log_root, ckpt_root, name, audio_text_weight, tool_weight, args): + run_log = os.path.join(log_root, f"{name}.log") + run_ckpt = os.path.join(ckpt_root, name) + if not args.no_checkpoints: + os.makedirs(run_ckpt, exist_ok=True) + + cmd = [ + _needle_path(repo_root), + "train", + "--epochs", str(args.epochs), + "--eval-every", str(args.eval_every), + "--audio-text-contrastive-weight", str(audio_text_weight), + "--tool-contrastive-weight", str(tool_weight), + ] + if args.no_checkpoints: + cmd.append("--no-checkpoints") + else: + cmd.extend(["--checkpoint-dir", run_ckpt]) + if args.skip_epoch_extras: + cmd.append("--skip-epoch-extras") + if args.cfg_inference: + cmd.append("--cfg-inference") + + with open(run_log, "w") as f: + f.write("COMMAND: " + " ".join(cmd) + "\n\n") + f.flush() + proc = subprocess.run(cmd, cwd=repo_root, stdout=f, stderr=subprocess.STDOUT, check=False) + return proc.returncode, run_log + + +def main(): + parser = argparse.ArgumentParser(description="Run contrastive-weight sweeps with separate log files") + parser.add_argument("--mode", choices=["audio5", "grid25"], default="audio5") + parser.add_argument("--epochs", type=int, default=1) + parser.add_argument("--eval-every", type=int, default=1000000) + parser.add_argument("--skip-epoch-extras", action="store_true") + parser.add_argument("--no-checkpoints", action="store_true") + parser.add_argument("--cfg-inference", action="store_true") + parser.add_argument("--log-root", type=str, default="logs/contrastive_sweeps") + parser.add_argument("--checkpoint-root", type=str, default="checkpoints/contrastive_sweeps") + args = parser.parse_args() + + repo_root = os.path.dirname(os.path.dirname(__file__)) + stamp = datetime.utcnow().strftime("%Y%m%d_%H%M%S") + run_root = os.path.join(repo_root, args.log_root, f"{args.mode}_{stamp}") + os.makedirs(run_root, exist_ok=True) + ckpt_root = os.path.join(repo_root, args.checkpoint_root, f"{args.mode}_{stamp}") + if not args.no_checkpoints: + os.makedirs(ckpt_root, exist_ok=True) + + combos = [] + if args.mode == "audio5": + for audio_text_weight in WEIGHTS: + combos.append((audio_text_weight, 0.0)) + else: + for audio_text_weight in WEIGHTS: + for tool_weight in WEIGHTS: + combos.append((audio_text_weight, tool_weight)) + + summary_path = os.path.join(run_root, "summary.txt") + with open(summary_path, "w") as summary: + summary.write(f"mode={args.mode}\n") + summary.write(f"epochs={args.epochs}\n") + summary.write(f"eval_every={args.eval_every}\n") + summary.write(f"skip_epoch_extras={args.skip_epoch_extras}\n") + summary.write(f"no_checkpoints={args.no_checkpoints}\n") + summary.write(f"cfg_inference={args.cfg_inference}\n\n") + for audio_text_weight, tool_weight in combos: + name = f"at_{_fmt_weight(audio_text_weight)}__tool_{_fmt_weight(tool_weight)}" + code, run_log = _run_one(repo_root, run_root, ckpt_root, name, audio_text_weight, tool_weight, args) + summary.write(f"{name}\treturncode={code}\tlog={run_log}\n") + summary.flush() + print(f"{name}: returncode={code} log={run_log}") + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_sweep_logs.py b/scripts/summarize_sweep_logs.py new file mode 100644 index 0000000..3d601fd --- /dev/null +++ b/scripts/summarize_sweep_logs.py @@ -0,0 +1,53 @@ +import argparse +import csv +import re +from pathlib import Path + + +PATTERNS = { + "text_loss": re.compile(r"Text loss\s+([0-9.]+)"), + "text_val_ppl": re.compile(r"Text val ppl\s+([0-9.]+)"), + "speech_loss": re.compile(r"Speech loss\s+([0-9.]+)"), + "speech_val_ppl": re.compile(r"Speech val ppl\s+([0-9.]+)"), +} + + +def _extract_last(text, pattern): + matches = pattern.findall(text) + return matches[-1] if matches else "" + + +def _rows(log_dir): + for log_path in sorted(Path(log_dir).glob("*.log")): + text = log_path.read_text(errors="ignore") + row = {"run": log_path.stem} + for key, pattern in PATTERNS.items(): + row[key] = _extract_last(text, pattern) + yield row + + +def main(): + parser = argparse.ArgumentParser(description="Extract final epoch metrics from sweep logs") + parser.add_argument("log_dir", type=str, help="Directory containing per-run .log files") + parser.add_argument("--csv", action="store_true", help="Print CSV instead of TSV") + args = parser.parse_args() + + fieldnames = ["run", "text_loss", "text_val_ppl", "speech_loss", "speech_val_ppl"] + rows = sorted(list(_rows(args.log_dir)), key=lambda x: float(x['speech_val_ppl'])) + + if args.csv: + writer = csv.DictWriter( + open("/dev/stdout", "w", newline=""), + fieldnames=fieldnames, + ) + writer.writeheader() + writer.writerows(rows) + return + + print("\t".join(fieldnames)) + for row in rows: + print("\t".join(row[name] for name in fieldnames)) + + +if __name__ == "__main__": + main() diff --git a/src/cli.py b/src/cli.py index 7a6cbcb..4f50b25 100644 --- a/src/cli.py +++ b/src/cli.py @@ -64,6 +64,8 @@ def main(): help="Auxiliary paired audio-text SigLIP loss weight") p.add_argument("--skip-epoch-extras", action="store_true", help="Skip epoch-end throughput benchmark and qualitative sample generation") + p.add_argument("--no-checkpoints", action="store_true", + help="Skip epoch checkpoint writes") p = sub.add_parser("tokenize", add_help=False) p.add_argument("--max-samples", type=int, default=None, diff --git a/src/train.py b/src/train.py index 3db8598..d5cecfc 100644 --- a/src/train.py +++ b/src/train.py @@ -1014,12 +1014,14 @@ def _tile_sp(arr): near_zero = sum(int(np.sum(np.abs(x) < 1e-6)) for x in jax.tree.leaves(params_np)) sparsity = near_zero / total_params * 100 - ckpt_name = f"needle_{args.num_layers}_{args.d_model}_{global_step}.pkl" - ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) - with open(ckpt_path, "wb") as f: - pickle.dump({"params": params_np, "config": config.__dict__}, f) + ckpt_path = "(skipped)" + if not getattr(args, "no_checkpoints", False): + ckpt_name = f"needle_{args.num_layers}_{args.d_model}_{global_step}.pkl" + ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) + with open(ckpt_path, "wb") as f: + pickle.dump({"params": params_np, "config": config.__dict__}, f) + print(f" [epoch {epoch + 1}] checkpoint saved in {time.perf_counter() - epoch_eval_t0:.1f}s") del params_np - print(f" [epoch {epoch + 1}] checkpoint saved in {time.perf_counter() - epoch_eval_t0:.1f}s") tp = {"tokens_per_second": float("nan"), "avg_latency_s": float("nan")} unified_samples = [] From f35f1ea720ea13eb730647bad57deb55aefe0c87 Mon Sep 17 00:00:00 2001 From: Karen Mosoyan Date: Wed, 11 Mar 2026 10:30:21 +0000 Subject: [PATCH 5/6] cleaned up diff --- .../grid25_20260310_072111/summary.txt | 6 - .../grid25_20260310_072312/summary.txt | 32 -- scripts/run_contrastive_sweep.py | 97 ---- scripts/summarize_sweep_logs.py | 53 -- src/cli.py | 27 +- src/pretrain.py | 76 +-- src/run.py | 49 +- src/test.py | 65 +-- src/tokenize_data.py | 14 - src/tool_cfg.py | 526 ------------------ src/train.py | 8 +- 11 files changed, 46 insertions(+), 907 deletions(-) delete mode 100644 logs/contrastive_sweeps/grid25_20260310_072111/summary.txt delete mode 100644 logs/contrastive_sweeps/grid25_20260310_072312/summary.txt delete mode 100644 scripts/run_contrastive_sweep.py delete mode 100644 scripts/summarize_sweep_logs.py delete mode 100644 src/tool_cfg.py diff --git a/logs/contrastive_sweeps/grid25_20260310_072111/summary.txt b/logs/contrastive_sweeps/grid25_20260310_072111/summary.txt deleted file mode 100644 index e1d821f..0000000 --- a/logs/contrastive_sweeps/grid25_20260310_072111/summary.txt +++ /dev/null @@ -1,6 +0,0 @@ -mode=grid25 -epochs=1 -eval_every=1000000 -skip_epoch_extras=True -cfg_inference=False - diff --git a/logs/contrastive_sweeps/grid25_20260310_072312/summary.txt b/logs/contrastive_sweeps/grid25_20260310_072312/summary.txt deleted file mode 100644 index ac0b234..0000000 --- a/logs/contrastive_sweeps/grid25_20260310_072312/summary.txt +++ /dev/null @@ -1,32 +0,0 @@ -mode=grid25 -epochs=1 -eval_every=1000000 -skip_epoch_extras=True -no_checkpoints=True -cfg_inference=False - -at_0p0__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_0p0.log -at_0p0__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_0p1.log -at_0p0__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_0p5.log -at_0p0__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_1p0.log -at_0p0__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p0__tool_2p0.log -at_0p1__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_0p0.log -at_0p1__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_0p1.log -at_0p1__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_0p5.log -at_0p1__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_1p0.log -at_0p1__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p1__tool_2p0.log -at_0p5__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_0p0.log -at_0p5__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_0p1.log -at_0p5__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_0p5.log -at_0p5__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_1p0.log -at_0p5__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_0p5__tool_2p0.log -at_1p0__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_0p0.log -at_1p0__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_0p1.log -at_1p0__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_0p5.log -at_1p0__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_1p0.log -at_1p0__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_1p0__tool_2p0.log -at_2p0__tool_0p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_0p0.log -at_2p0__tool_0p1 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_0p1.log -at_2p0__tool_0p5 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_0p5.log -at_2p0__tool_1p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_1p0.log -at_2p0__tool_2p0 returncode=0 log=/home/karen/needle/logs/contrastive_sweeps/grid25_20260310_072312/at_2p0__tool_2p0.log diff --git a/scripts/run_contrastive_sweep.py b/scripts/run_contrastive_sweep.py deleted file mode 100644 index c087ab8..0000000 --- a/scripts/run_contrastive_sweep.py +++ /dev/null @@ -1,97 +0,0 @@ -import argparse -import os -import subprocess -import sys -from datetime import datetime - - -WEIGHTS = [0.0, 0.1, 0.5, 1.0, 2.0] - - -def _fmt_weight(value): - text = str(value) - return text.replace(".", "p") - - -def _needle_path(repo_root): - local = os.path.join(repo_root, ".venv", "bin", "needle") - return local if os.path.exists(local) else "needle" - - -def _run_one(repo_root, log_root, ckpt_root, name, audio_text_weight, tool_weight, args): - run_log = os.path.join(log_root, f"{name}.log") - run_ckpt = os.path.join(ckpt_root, name) - if not args.no_checkpoints: - os.makedirs(run_ckpt, exist_ok=True) - - cmd = [ - _needle_path(repo_root), - "train", - "--epochs", str(args.epochs), - "--eval-every", str(args.eval_every), - "--audio-text-contrastive-weight", str(audio_text_weight), - "--tool-contrastive-weight", str(tool_weight), - ] - if args.no_checkpoints: - cmd.append("--no-checkpoints") - else: - cmd.extend(["--checkpoint-dir", run_ckpt]) - if args.skip_epoch_extras: - cmd.append("--skip-epoch-extras") - if args.cfg_inference: - cmd.append("--cfg-inference") - - with open(run_log, "w") as f: - f.write("COMMAND: " + " ".join(cmd) + "\n\n") - f.flush() - proc = subprocess.run(cmd, cwd=repo_root, stdout=f, stderr=subprocess.STDOUT, check=False) - return proc.returncode, run_log - - -def main(): - parser = argparse.ArgumentParser(description="Run contrastive-weight sweeps with separate log files") - parser.add_argument("--mode", choices=["audio5", "grid25"], default="audio5") - parser.add_argument("--epochs", type=int, default=1) - parser.add_argument("--eval-every", type=int, default=1000000) - parser.add_argument("--skip-epoch-extras", action="store_true") - parser.add_argument("--no-checkpoints", action="store_true") - parser.add_argument("--cfg-inference", action="store_true") - parser.add_argument("--log-root", type=str, default="logs/contrastive_sweeps") - parser.add_argument("--checkpoint-root", type=str, default="checkpoints/contrastive_sweeps") - args = parser.parse_args() - - repo_root = os.path.dirname(os.path.dirname(__file__)) - stamp = datetime.utcnow().strftime("%Y%m%d_%H%M%S") - run_root = os.path.join(repo_root, args.log_root, f"{args.mode}_{stamp}") - os.makedirs(run_root, exist_ok=True) - ckpt_root = os.path.join(repo_root, args.checkpoint_root, f"{args.mode}_{stamp}") - if not args.no_checkpoints: - os.makedirs(ckpt_root, exist_ok=True) - - combos = [] - if args.mode == "audio5": - for audio_text_weight in WEIGHTS: - combos.append((audio_text_weight, 0.0)) - else: - for audio_text_weight in WEIGHTS: - for tool_weight in WEIGHTS: - combos.append((audio_text_weight, tool_weight)) - - summary_path = os.path.join(run_root, "summary.txt") - with open(summary_path, "w") as summary: - summary.write(f"mode={args.mode}\n") - summary.write(f"epochs={args.epochs}\n") - summary.write(f"eval_every={args.eval_every}\n") - summary.write(f"skip_epoch_extras={args.skip_epoch_extras}\n") - summary.write(f"no_checkpoints={args.no_checkpoints}\n") - summary.write(f"cfg_inference={args.cfg_inference}\n\n") - for audio_text_weight, tool_weight in combos: - name = f"at_{_fmt_weight(audio_text_weight)}__tool_{_fmt_weight(tool_weight)}" - code, run_log = _run_one(repo_root, run_root, ckpt_root, name, audio_text_weight, tool_weight, args) - summary.write(f"{name}\treturncode={code}\tlog={run_log}\n") - summary.flush() - print(f"{name}: returncode={code} log={run_log}") - - -if __name__ == "__main__": - main() diff --git a/scripts/summarize_sweep_logs.py b/scripts/summarize_sweep_logs.py deleted file mode 100644 index 3d601fd..0000000 --- a/scripts/summarize_sweep_logs.py +++ /dev/null @@ -1,53 +0,0 @@ -import argparse -import csv -import re -from pathlib import Path - - -PATTERNS = { - "text_loss": re.compile(r"Text loss\s+([0-9.]+)"), - "text_val_ppl": re.compile(r"Text val ppl\s+([0-9.]+)"), - "speech_loss": re.compile(r"Speech loss\s+([0-9.]+)"), - "speech_val_ppl": re.compile(r"Speech val ppl\s+([0-9.]+)"), -} - - -def _extract_last(text, pattern): - matches = pattern.findall(text) - return matches[-1] if matches else "" - - -def _rows(log_dir): - for log_path in sorted(Path(log_dir).glob("*.log")): - text = log_path.read_text(errors="ignore") - row = {"run": log_path.stem} - for key, pattern in PATTERNS.items(): - row[key] = _extract_last(text, pattern) - yield row - - -def main(): - parser = argparse.ArgumentParser(description="Extract final epoch metrics from sweep logs") - parser.add_argument("log_dir", type=str, help="Directory containing per-run .log files") - parser.add_argument("--csv", action="store_true", help="Print CSV instead of TSV") - args = parser.parse_args() - - fieldnames = ["run", "text_loss", "text_val_ppl", "speech_loss", "speech_val_ppl"] - rows = sorted(list(_rows(args.log_dir)), key=lambda x: float(x['speech_val_ppl'])) - - if args.csv: - writer = csv.DictWriter( - open("/dev/stdout", "w", newline=""), - fieldnames=fieldnames, - ) - writer.writeheader() - writer.writerows(rows) - return - - print("\t".join(fieldnames)) - for row in rows: - print("\t".join(row[name] for name in fieldnames)) - - -if __name__ == "__main__": - main() diff --git a/src/cli.py b/src/cli.py index fdb0344..1bde625 100644 --- a/src/cli.py +++ b/src/cli.py @@ -33,17 +33,29 @@ def main(): p.add_argument("--seed", type=int, default=42) p.add_argument("--eval-every", type=int, default=1000) p.add_argument("--max-eval-samples", type=int, default=None) + p.add_argument("--sparsity-ratio", type=float, default=0.0) p.add_argument("--group-size", type=int, default=32) + p.add_argument("--prune-interval", type=int, default=100, + help="Steps between mask updates during gradual pruning (default: 100)") + p.add_argument("--prune-start-frac", type=float, default=0.33, + help="Fraction of epoch to train before starting gradual pruning (default: 0.33)") + p.add_argument("--prune-end-frac", type=float, default=0.67, + help="Fraction of epoch at which pruning finishes and mask locks (default: 0.67)") p.add_argument("--activation", type=str, default="drelu", choices=["drelu", "swiglu", "geglu"]) p.add_argument("--num-memory-slots", type=int, default=64) + p.add_argument("--mat-factors", type=int, nargs="*", default=[2, 4, 8], + help="Matryoshka FFN shrink factors, e.g. 2=half width (default: 2 4 8)") + p.add_argument("--mat-shared-input", action="store_true", + help="Each unique input is repeated across all mat widths (default: unique input per width)") p.add_argument("--dropout", type=float, default=0.0, help="Dropout rate for residual connections (default: 0.1)") + p.add_argument("--no-speech", action="store_true", help="Disable speech training (text-only)") p.add_argument("--max-mel-len", type=int, default=1024, help="Max mel spectrogram frames (default: 1024)") p.add_argument("--n-mels", type=int, default=80, help="Number of mel frequency bins (default: 80)") - p.add_argument("--no-checkpoints", action="store_true", - help="Skip epoch checkpoint writes") + p.add_argument("--max-speech-samples", type=int, default=None, + help="Max voice-tool-call training samples (default: all)") p = sub.add_parser("pretrain", add_help=False) p.add_argument("--full", action="store_true") @@ -85,14 +97,16 @@ def main(): help="Toucan query-description contrastive loss weight") p.add_argument("--audio-text-contrastive-weight", type=float, default=1.0, help="Paired audio-text SigLIP loss weight") - p.add_argument("--no-checkpoints", action="store_true", - help="Skip epoch checkpoint writes") p = sub.add_parser("tokenize", add_help=False) p.add_argument("--max-samples", type=int, default=None, help="Limit samples per split (for dev/test)") p.add_argument("--cleanup", action="store_true", help="Delete local .data_cache/ after GCS upload") + p.add_argument("--n-mels", type=int, default=80, + help="Number of mel frequency bins (default: 80)") + p.add_argument("--max-mel-len", type=int, default=1024, + help="Max mel spectrogram frames (default: 1024)") p.add_argument("--max-enc-len", type=int, default=256, help="Max encoder sequence length (default: 256)") p.add_argument("--max-dec-len", type=int, default=1024, @@ -111,8 +125,6 @@ def main(): help="Toucan subset to parse and cache during tokenization") p.add_argument("--toucan-max-samples", type=int, default=None, help="Optional max samples for Toucan Stage 1 preprocessing") - p.add_argument("--clear-gcs-cache", action="store_true", - help="Delete shared prepared-cache artifacts in GCS before uploading new ones") p.add_argument("--overwrite-gcs-tokenizer", action="store_true", help="Overwrite the shared GCS tokenizer after retraining locally") @@ -123,7 +135,6 @@ def main(): p.add_argument("--audio", type=str, nargs="*", help="Audio file paths for voice-to-tool-call") p.add_argument("--max-len", type=int, default=512) p.add_argument("--seed", type=int, default=0) - p.add_argument("--cfg-inference", action="store_true") p = sub.add_parser("test", add_help=False) p.add_argument("--checkpoint", type=str, required=True) @@ -137,7 +148,6 @@ def main(): p.add_argument("--voice-tc-samples", type=int, default=50, help="Samples for voice-to-tool-call eval (default: 50)") p.add_argument("--throughput-runs", type=int, default=10) - p.add_argument("--cfg-inference", action="store_true") p = sub.add_parser("evaluate", add_help=False) p.add_argument("--checkpoint", type=str, required=True) @@ -203,6 +213,7 @@ def main(): args.num_layers = 12 args.num_dec_layers = 4 args.num_memory_slots = 128 + args.mat_factors = [2, 3, 4, 8, 16] from .train import train train(args) elif args.command == "run": diff --git a/src/pretrain.py b/src/pretrain.py index b4fc591..1c76a98 100644 --- a/src/pretrain.py +++ b/src/pretrain.py @@ -38,45 +38,6 @@ _AUDIO_TEXT_CONTRASTIVE_WEIGHT = 1.0 -def _assert_stage1_batch_shapes(mel, transcript_text, tgt_in, tgt_out, loss_mask, - tc_q, tc_tools, tc_labels, tc_tool_mask, - effective_batch_size, max_mel_len, n_mels, - max_enc_len, max_dec_len): - expected = { - "mel": (effective_batch_size, max_mel_len, n_mels), - "transcript_text": (effective_batch_size, max_enc_len), - "tgt_in": (effective_batch_size, max_dec_len), - "tgt_out": (effective_batch_size, max_dec_len), - "loss_mask": (effective_batch_size, max_dec_len), - "tc_q": (effective_batch_size, max_enc_len), - } - actual = { - "mel": mel.shape, - "transcript_text": transcript_text.shape, - "tgt_in": tgt_in.shape, - "tgt_out": tgt_out.shape, - "loss_mask": loss_mask.shape, - "tc_q": tc_q.shape, - } - for name, shape in expected.items(): - if actual[name] != shape: - raise ValueError(f"Unexpected {name} shape {actual[name]}, expected {shape}") - - if tc_tools.shape[0] != effective_batch_size or tc_tools.shape[2] != max_enc_len: - raise ValueError( - f"Unexpected tc_tools shape {tc_tools.shape}, expected " - f"({effective_batch_size}, n_tools, {max_enc_len})" - ) - if tc_labels.shape != tc_tool_mask.shape or tc_labels.shape[0] != effective_batch_size: - raise ValueError( - f"Unexpected Toucan label/mask shapes {tc_labels.shape} and {tc_tool_mask.shape}" - ) - if tc_labels.shape[1] != tc_tools.shape[1]: - raise ValueError( - f"Toucan tool dimension mismatch: tc_tools={tc_tools.shape}, tc_labels={tc_labels.shape}" - ) - - def _tool_contrastive_loss(state, params, contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask): pad_id = 0 q_mask = make_padding_mask(contrastive_q, pad_id) @@ -321,10 +282,6 @@ def pretrain(args): ckpt_params = None effective_batch_size = args.batch_size * num_devices - if effective_batch_size % num_devices != 0: - raise ValueError( - f"Effective batch size {effective_batch_size} must be divisible by num_devices {num_devices}" - ) batches_per_epoch = count_batches(len(train_prepared["text_inputs"]), effective_batch_size) total_steps = batches_per_epoch * args.epochs warmup_steps = max(1, int(total_steps * args.warmup_ratio)) @@ -366,7 +323,6 @@ def pretrain(args): global_step = 0 last_val_ppl = None eval_every = getattr(args, "eval_every", 1000) - checked_batch_shapes = False for epoch in range(args.epochs): losses = [] @@ -391,30 +347,6 @@ def pretrain(args): mel, text_enc, tgt_in, tgt_out, lm = next(speech_iter) tc_q, tc_tools, tc_labels, tc_tool_mask = next(contrastive_iter) - if not checked_batch_shapes: - _assert_stage1_batch_shapes( - mel, - text_enc, - tgt_in, - tgt_out, - lm, - tc_q, - tc_tools, - tc_labels, - tc_tool_mask, - effective_batch_size, - args.max_mel_len, - args.n_mels, - args.max_enc_len, - args.max_dec_len, - ) - print( - " Stage 1 batch shapes " - f"mel={mel.shape} text={text_enc.shape} tgt={tgt_in.shape} " - f"tc_q={tc_q.shape} tc_tools={tc_tools.shape}" - ) - checked_batch_shapes = True - mel_b = shard_batch(mel, num_devices) text_enc_b = shard_batch(text_enc, num_devices) tgt_in_b = shard_batch(tgt_in, num_devices) @@ -496,11 +428,9 @@ def pretrain(args): args.max_dec_len, max_eval_samples=getattr(args, "max_eval_samples", None), ) - ckpt_path = "(skipped)" - if not getattr(args, "no_checkpoints", False): - ckpt_name = f"needle_stage1_{config.num_encoder_layers}_{config.d_model}_{global_step}.pkl" - ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) - save_checkpoint(ckpt_path, eval_params, config, extra={"stage": "stage1"}) + ckpt_name = f"needle_stage1_{config.num_encoder_layers}_{config.d_model}_{global_step}.pkl" + ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) + save_checkpoint(ckpt_path, eval_params, config, extra={"stage": "stage1"}) print(f"\n Epoch {epoch + 1}/{args.epochs}") print(f" Train loss {sum(losses) / max(len(losses), 1):.4f}") diff --git a/src/run.py b/src/run.py index fa3458a..f029414 100644 --- a/src/run.py +++ b/src/run.py @@ -14,7 +14,6 @@ make_padding_mask, make_mel_padding_mask, ) -from .tool_cfg import ToolCallCFG _decode_fn_cache = {} @@ -49,17 +48,7 @@ def load_checkpoint(path): return params, config -def _select_next_token(next_logits, allowed_ids=None): - if allowed_ids is None: - return int(jnp.argmax(next_logits)) - allowed_ids = jnp.array(allowed_ids, dtype=jnp.int32) - masked = jnp.full_like(next_logits, jnp.finfo(next_logits.dtype).min) - masked = masked.at[allowed_ids].set(next_logits[allowed_ids]) - return int(jnp.argmax(masked)) - - -def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed=0, - stream=True, task_token_id=None, use_cfg=False): +def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed=0, stream=True, task_token_id=None): """Generate tool-call output. Encoder: query only. @@ -77,7 +66,7 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= {"params": params}, enc_input, src_mask=src_mask, method="encode" ) - tools_tokens = tokenizer.encode_tool_schema(tools) + tools_tokens = tokenizer.encode(tools) prefix = [eos_id, tool_call_id] + tools_tokens prefix_len = min(len(prefix), max_gen_len) @@ -88,8 +77,6 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= decode_fn = _get_decode_fn(model, max_gen_len) generated_tokens = [] - cfg = ToolCallCFG(tokenizer, tools) if use_cfg else None - streamed_text = "" if stream: sys.stdout.write(f"\n") @@ -99,21 +86,16 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= for i in range(prefix_len - 1, max_gen_len - 1): next_logits = logits[0, i] - allowed_ids = cfg.allowed_next_ids(training=False) if cfg is not None else None - next_token = _select_next_token(next_logits, allowed_ids) + next_token = int(jnp.argmax(next_logits)) if next_token == eos_id: break generated_tokens.append(next_token) - if cfg is not None: - cfg.step(next_token) dec_buffer = dec_buffer.at[0, i + 1].set(next_token) if stream: - current_text = tokenizer.decode_structured(generated_tokens) - sys.stdout.write(current_text[len(streamed_text):]) - streamed_text = current_text + sys.stdout.write(tokenizer.decode([next_token])) sys.stdout.flush() logits = decode_fn(params, dec_buffer, encoder_out) @@ -121,7 +103,7 @@ def generate(model, params, tokenizer, query, tools="[]", max_gen_len=512, seed= if stream: sys.stdout.write("\n") - return tokenizer.decode_structured(generated_tokens) + return tokenizer.decode(generated_tokens) def load_audio(path, target_sr=16000): @@ -138,8 +120,7 @@ def load_audio(path, target_sr=16000): return audio, sr -def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools="[]", - max_gen_len=512, seed=0, stream=True, use_cfg=False): +def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools="[]", max_gen_len=512, seed=0, stream=True): """Generate tool-call output from audio using the speech encoder pathway. mel -> encode_speech -> decoder [BOS, , tools_tokens...] -> greedy decode. @@ -158,7 +139,7 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" {"params": params}, mel_input, src_mask=src_mask, deterministic=True, method="encode_speech" ) - tools_tokens = tokenizer.encode_tool_schema(tools) + tools_tokens = tokenizer.encode(tools) prefix = [eos_id, tool_call_id] + tools_tokens prefix_len = min(len(prefix), max_gen_len) @@ -169,8 +150,6 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" decode_fn = _get_decode_fn(model, max_gen_len) generated_tokens = [] - cfg = ToolCallCFG(tokenizer, tools) if use_cfg else None - streamed_text = "" if stream: sys.stdout.write("\n") @@ -180,21 +159,16 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" for i in range(prefix_len - 1, max_gen_len - 1): next_logits = logits[0, i] - allowed_ids = cfg.allowed_next_ids(training=False) if cfg is not None else None - next_token = _select_next_token(next_logits, allowed_ids) + next_token = int(jnp.argmax(next_logits)) if next_token == eos_id: break generated_tokens.append(next_token) - if cfg is not None: - cfg.step(next_token) dec_buffer = dec_buffer.at[0, i + 1].set(next_token) if stream: - current_text = tokenizer.decode_structured(generated_tokens) - sys.stdout.write(current_text[len(streamed_text):]) - streamed_text = current_text + sys.stdout.write(tokenizer.decode([next_token])) sys.stdout.flush() logits = decode_fn(params, dec_buffer, encoder_out) @@ -202,7 +176,7 @@ def generate_from_audio(model, params, tokenizer, audio_array, sr=16000, tools=" if stream: sys.stdout.write("\n") - return tokenizer.decode_structured(generated_tokens) + return tokenizer.decode(generated_tokens) def main(args): @@ -232,7 +206,6 @@ def main(args): max_gen_len=args.max_len, seed=args.seed + i, stream=True, - use_cfg=getattr(args, "cfg_inference", False), ) return @@ -260,7 +233,6 @@ def main(args): max_gen_len=args.max_len, seed=args.seed + i, stream=True, - use_cfg=getattr(args, "cfg_inference", False), ) @@ -272,7 +244,6 @@ def parse_args(): parser.add_argument("--audio", type=str, nargs="*", help="Audio file paths for voice-to-tool-call") parser.add_argument("--max-len", type=int, default=512) parser.add_argument("--seed", type=int, default=0) - parser.add_argument("--cfg-inference", action="store_true") return parser.parse_args() diff --git a/src/test.py b/src/test.py index 857157e..93502d2 100644 --- a/src/test.py +++ b/src/test.py @@ -7,15 +7,7 @@ import numpy as np import optax -from .data import ( - _load_cache_metadata, - get_batches, - get_tokenizer, - load_tool_calls, - prepare_tool_call_pairs, - load_tool_call_audio, - load_example_with_audio, -) +from .data import get_batches, get_tokenizer, load_tool_calls, prepare_tool_call_pairs, load_tool_call_audio, load_example_with_audio from .model import ( EncoderDecoderTransformer, TransformerConfig, @@ -131,7 +123,7 @@ def benchmark_generation_quality(model, params, tokenizer, prompts, max_gen_len= generations = [] for i, prompt in enumerate(prompts): - text = generate(model, params, tokenizer, prompt, max_gen_len=max_gen_len, seed=i, stream=False) + text = generate(model, params, tokenizer, prompt, max_gen_len, temperature, seed=i, stream=False) generations.append(text) lengths = [len(tokenizer.encode(t)) for t in generations] @@ -176,19 +168,13 @@ def compute_wer(hypotheses, references): return total_edits / max(total_ref_words, 1) -def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len=512, - shuffle_before_split=False, shuffle_seed=42, use_cfg=False): +def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len=512): """Generate tool-call predictions and compute structured metrics.""" import json from .run import generate from .data import load_tool_calls - ds = load_tool_calls( - "validation", - max_samples=num_samples, - shuffle_before_split=shuffle_before_split, - shuffle_seed=shuffle_seed, - ) + ds = load_tool_calls("validation", max_samples=num_samples) total = 0 exact_match = 0 @@ -209,7 +195,7 @@ def benchmark_tool_calls(model, params, tokenizer, num_samples=200, max_gen_len= pred_text = generate( model, params, tokenizer, ex["query"], - tools=ex["tools"], max_gen_len=max_gen_len, seed=i, stream=False, use_cfg=use_cfg, + tools=ex["tools"], max_gen_len=max_gen_len, seed=i, stream=False, ).strip() try: @@ -282,18 +268,12 @@ def call_key(c): } -def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_gen_len=512, - shuffle_before_split=False, shuffle_seed=42, use_cfg=False): +def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_gen_len=512): """Generate tool-call predictions from audio and compute structured metrics.""" import json from .run import generate_from_audio - indices = load_tool_call_audio( - "validation", - max_samples=num_samples, - shuffle_before_split=shuffle_before_split, - shuffle_seed=shuffle_seed, - ) + indices = load_tool_call_audio("validation", max_samples=num_samples) total = 0 exact_match = 0 @@ -316,7 +296,7 @@ def benchmark_voice_tool_calls(model, params, tokenizer, num_samples=100, max_ge pred_text = generate_from_audio( model, params, tokenizer, pair["audio_array"], sr=pair["sampling_rate"], - tools=pair["tools"], max_gen_len=max_gen_len, seed=i, stream=False, use_cfg=use_cfg, + tools=pair["tools"], max_gen_len=max_gen_len, seed=i, stream=False, ).strip() try: @@ -393,9 +373,6 @@ def main(args): params, config = load_checkpoint(args.checkpoint) model = EncoderDecoderTransformer(config) tokenizer = get_tokenizer() - val_meta = _load_cache_metadata("val") or {} - split_shuffle = val_meta.get("shuffle_before_split", False) - split_seed = val_meta.get("split_seed", 42) param_count = sum(x.size for x in jax.tree.leaves(params)) print(f"\ncheckpoint: {args.checkpoint}") @@ -403,12 +380,7 @@ def main(args): print(f"config: d={config.d_model}, heads={config.num_heads}, layers={config.num_encoder_layers}/{config.num_decoder_layers}") print(f"\nevaluating tool-call perplexity ({args.max_eval_samples} samples)...") - ds = load_tool_calls( - "validation", - max_samples=args.max_eval_samples, - shuffle_before_split=split_shuffle, - shuffle_seed=split_seed, - ) + ds = load_tool_calls("validation", max_samples=args.max_eval_samples) enc_inputs, dec_inputs, dec_targets, loss_mask_arr, _ = prepare_tool_call_pairs( ds, tokenizer, max_enc_len=args.max_enc_len, max_dec_len=args.max_dec_len ) @@ -418,14 +390,7 @@ def main(args): tc = None if tc_samples > 0: print(f"\nevaluating tool-call accuracy ({tc_samples} samples)...") - tc = benchmark_tool_calls( - model, params, tokenizer, - num_samples=tc_samples, - max_gen_len=args.max_gen_len, - shuffle_before_split=split_shuffle, - shuffle_seed=split_seed, - use_cfg=getattr(args, "cfg_inference", False), - ) + tc = benchmark_tool_calls(model, params, tokenizer, num_samples=tc_samples, max_gen_len=args.max_gen_len) print(f"\n ─────────────────────────────────────") print(f" Tool-Call Metrics") @@ -455,14 +420,7 @@ def main(args): voice_tc_samples = getattr(args, "voice_tc_samples", 50) if voice_tc_samples > 0: print(f"\nevaluating voice-to-tool-call ({voice_tc_samples} samples)...") - vtc = benchmark_voice_tool_calls( - model, params, tokenizer, - num_samples=voice_tc_samples, - max_gen_len=args.max_gen_len, - shuffle_before_split=split_shuffle, - shuffle_seed=split_seed, - use_cfg=getattr(args, "cfg_inference", False), - ) + vtc = benchmark_voice_tool_calls(model, params, tokenizer, num_samples=voice_tc_samples, max_gen_len=args.max_gen_len) print(f"\n ─── Voice-Tool-Call Metrics ─────────") print(f" JSON parse rate {vtc['json_parse_rate']:>10.1%}") print(f" Exact match {vtc['exact_match']:>10.1%}") @@ -504,7 +462,6 @@ def parse_args(): help="Samples for tool-call accuracy eval (default: 200)") parser.add_argument("--voice-tc-samples", type=int, default=50, help="Samples for voice-to-tool-call eval (default: 50)") - parser.add_argument("--cfg-inference", action="store_true") return parser.parse_args() diff --git a/src/tokenize_data.py b/src/tokenize_data.py index 174268d..5484bca 100644 --- a/src/tokenize_data.py +++ b/src/tokenize_data.py @@ -18,7 +18,6 @@ from .data import ( CACHE_DIR, - GCS_CACHE_PATH, GCS_TOKENIZER_PATH, EMILIA_SPEECH_GCS_PREFIX, TOKENIZER_DIR, @@ -37,16 +36,6 @@ from .toucan import cache_toucan_examples -def _clear_gcs_cache(): - """Remove existing shared prepared-cache files.""" - path = GCS_CACHE_PATH + "/*" - print(f"Clearing {path} ...") - subprocess.run( - ["gcloud", "storage", "rm", "-r", path], - capture_output=True, text=True, - ) - - def _clear_gcs_tokenizer(): """Remove the shared tokenizer files.""" path = GCS_TOKENIZER_PATH + "*" @@ -68,9 +57,6 @@ def _clear_local_caches(): def tokenize(args): print("=== Clearing existing local caches ===") _clear_local_caches() - if getattr(args, "clear_gcs_cache", False): - print("\n=== Clearing shared GCS cache ===") - _clear_gcs_cache() if getattr(args, "overwrite_gcs_tokenizer", False): print("\n=== Clearing shared GCS tokenizer ===") _clear_gcs_tokenizer() diff --git a/src/tool_cfg.py b/src/tool_cfg.py deleted file mode 100644 index be68f6b..0000000 --- a/src/tool_cfg.py +++ /dev/null @@ -1,526 +0,0 @@ -import re - -import numpy as np - -from .data import ( - EOS_ID, - JSON_COLON, - JSON_COMMA, - JSON_FALSE, - JSON_KEY_ARGUMENTS, - JSON_KEY_NAME, - JSON_LBRACE, - JSON_LBRACK, - JSON_NULL, - JSON_QUOTE, - JSON_RBRACE, - JSON_RBRACK, - JSON_TRUE, - _normalize_tool_schema_spec, -) - - -_NUMBER_PREFIX_RE = re.compile(r"^-?(?:0|[1-9]\d*)(?:\.\d*)?(?:[eE][+-]?\d*)?$|^-?(?:0|[1-9]\d*)?$|^-?(?:0|[1-9]\d*)\.$|^-?(?:0|[1-9]\d*)(?:\.\d+)?[eE]?$|^-?(?:0|[1-9]\d*)(?:\.\d+)?[eE][+-]?$") -_NUMBER_COMPLETE_RE = re.compile(r"^-?(?:0|[1-9]\d*)(?:\.\d+)?(?:[eE][+-]?\d+)?$") - - -def _is_number_prefix(text): - return bool(text) and _NUMBER_PREFIX_RE.match(text) is not None - - -def _is_complete_number(text): - return bool(text) and _NUMBER_COMPLETE_RE.match(text) is not None - - -def _schema_type_name(value): - if isinstance(value, str): - lowered = value.lower() - if lowered in {"str", "string", "text"}: - return "string" - if lowered in {"bool", "boolean"}: - return "boolean" - if lowered in {"int", "integer", "long"}: - return "integer" - if lowered in {"float", "double", "number", "numeric"}: - return "number" - if lowered in {"list", "array"}: - return "array" - if lowered in {"dict", "map", "object"}: - return "object" - if isinstance(value, list): - return "array" - if isinstance(value, dict): - return "object" - return "any" - - -def _build_trie(sequences): - root = {} - for index, sequence in sequences: - node = root - for token_id in sequence: - node = node.setdefault(int(token_id), {}) - node["_end"] = index - return root - - -class JsonValueParser: - def __init__(self, tokenizer, type_hint="any"): - self.tokenizer = tokenizer - self.type_hint = type_hint - self.regular_ids = tokenizer.regular_token_ids - self.quote_id = tokenizer.json_token_ids[JSON_QUOTE] - self.lbrace_id = tokenizer.json_token_ids[JSON_LBRACE] - self.rbrace_id = tokenizer.json_token_ids[JSON_RBRACE] - self.lbrack_id = tokenizer.json_token_ids[JSON_LBRACK] - self.rbrack_id = tokenizer.json_token_ids[JSON_RBRACK] - self.colon_id = tokenizer.json_token_ids[JSON_COLON] - self.comma_id = tokenizer.json_token_ids[JSON_COMMA] - self.true_id = tokenizer.json_token_ids[JSON_TRUE] - self.false_id = tokenizer.json_token_ids[JSON_FALSE] - self.null_id = tokenizer.json_token_ids[JSON_NULL] - self.mode = "expect_value" - self.stack = [] - self.number_text = "" - self.complete = False - self._regular_surfaces = { - int(token_id): tokenizer.token_surface(int(token_id)) - for token_id in self.regular_ids - } - self._number_prefix_ids = np.array( - [token_id for token_id, text in self._regular_surfaces.items() if _is_number_prefix(text)], - dtype=np.int32, - ) - self._number_cache = {} - - def _start_value_ids(self, type_hint=None): - hint = type_hint or self.type_hint - if hint == "string": - return np.array([self.quote_id], dtype=np.int32) - if hint == "boolean": - return np.array([self.true_id, self.false_id], dtype=np.int32) - if hint in {"integer", "number"}: - return self._number_prefix_ids - if hint == "array": - return np.array([self.lbrack_id], dtype=np.int32) - if hint == "object": - return np.array([self.lbrace_id], dtype=np.int32) - return np.concatenate( - [ - np.array([self.quote_id, self.lbrace_id, self.lbrack_id, self.true_id, self.false_id, self.null_id], dtype=np.int32), - self._number_prefix_ids, - ] - ) - - def _number_allowed_ids(self, training=False): - if training: - return None - cached = self._number_cache.get(self.number_text) - if cached is not None: - return cached - allowed = [] - for token_id, surface in self._regular_surfaces.items(): - if _is_number_prefix(self.number_text + surface): - allowed.append(token_id) - if self._is_value_complete(): - allowed.extend(self._value_end_ids().tolist()) - cached = np.array(sorted(set(allowed)), dtype=np.int32) - self._number_cache[self.number_text] = cached - return cached - - def _value_end_ids(self): - if not self.stack: - return np.array([EOS_ID], dtype=np.int32) - frame = self.stack[-1] - if frame["kind"] == "array": - return np.array([self.comma_id, self.rbrack_id], dtype=np.int32) - return np.array([self.comma_id, self.rbrace_id], dtype=np.int32) - - def _value_finished(self): - if not self.stack: - self.complete = True - self.mode = "done" - return - frame = self.stack[-1] - if frame["kind"] == "array": - frame["mode"] = "after_value" - self.mode = "array_after_value" - else: - frame["mode"] = "after_value" - self.mode = "object_after_value" - - def allowed_next_ids(self, training=False): - if self.mode == "done": - return np.array([EOS_ID], dtype=np.int32) - if self.mode == "expect_value": - ids = self._start_value_ids() - if training and self.type_hint == "any": - return np.array([self.quote_id, self.lbrace_id, self.lbrack_id, self.true_id, self.false_id, self.null_id], dtype=np.int32) - return ids - if self.mode == "string": - if training: - return None - return np.concatenate([self.regular_ids, np.array([self.quote_id], dtype=np.int32)]) - if self.mode == "number": - return self._number_allowed_ids(training=training) - if self.mode == "array_value_or_end": - return np.concatenate([np.array([self.rbrack_id], dtype=np.int32), self._start_value_ids("any")]) - if self.mode == "array_after_value": - return np.array([self.comma_id, self.rbrack_id], dtype=np.int32) - if self.mode == "object_key_or_end": - return np.array([self.quote_id, self.rbrace_id], dtype=np.int32) - if self.mode == "object_key_string": - if training: - return None - return np.concatenate([self.regular_ids, np.array([self.quote_id], dtype=np.int32)]) - if self.mode == "object_after_key": - return np.array([self.colon_id], dtype=np.int32) - if self.mode == "object_after_value": - return np.array([self.comma_id, self.rbrace_id], dtype=np.int32) - return None - - def step(self, token_id): - token_id = int(token_id) - if self.mode == "expect_value": - if token_id == self.quote_id: - self.mode = "string" - return - if token_id == self.lbrack_id: - self.stack.append({"kind": "array", "mode": "value_or_end"}) - self.mode = "array_value_or_end" - return - if token_id == self.lbrace_id: - self.stack.append({"kind": "object", "mode": "key_or_end"}) - self.mode = "object_key_or_end" - return - if token_id in {self.true_id, self.false_id, self.null_id}: - self._value_finished() - return - self.mode = "number" - self.number_text = self._regular_surfaces.get(token_id, "") - return - - if self.mode == "string": - if token_id == self.quote_id: - self._value_finished() - return - - if self.mode == "number": - if token_id in set(self._value_end_ids().tolist()): - if token_id == EOS_ID: - self.complete = True - self.mode = "done" - return - frame = self.stack[-1] - if frame["kind"] == "array": - if token_id == self.comma_id: - frame["mode"] = "value" - self.mode = "expect_value" - self.type_hint = "any" - else: - self.stack.pop() - self._value_finished() - else: - if token_id == self.comma_id: - frame["mode"] = "key_or_end" - self.mode = "object_key_or_end" - else: - self.stack.pop() - self._value_finished() - self.number_text = "" - return - self.number_text += self._regular_surfaces.get(token_id, "") - return - - if self.mode == "array_value_or_end": - if token_id == self.rbrack_id: - self.stack.pop() - self._value_finished() - return - self.mode = "expect_value" - self.type_hint = "any" - self.step(token_id) - return - - if self.mode == "array_after_value": - frame = self.stack[-1] - if token_id == self.comma_id: - frame["mode"] = "value" - self.mode = "expect_value" - self.type_hint = "any" - return - self.stack.pop() - self._value_finished() - return - - if self.mode == "object_key_or_end": - if token_id == self.rbrace_id: - self.stack.pop() - self._value_finished() - return - self.mode = "object_key_string" - return - - if self.mode == "object_key_string": - if token_id == self.quote_id: - self.mode = "object_after_key" - return - - if self.mode == "object_after_key": - self.mode = "expect_value" - self.type_hint = "any" - return - - if self.mode == "object_after_value": - frame = self.stack[-1] - if token_id == self.comma_id: - frame["mode"] = "key_or_end" - self.mode = "object_key_or_end" - return - self.stack.pop() - self._value_finished() - - def _is_value_complete(self): - return _is_complete_number(self.number_text) - - -class ToolCallCFG: - def __init__(self, tokenizer, tools_text): - self.tokenizer = tokenizer - self.tools_text = tools_text - self.tools = _normalize_tool_schema_spec(tools_text) - self.tool_names = [tool["name"] for tool in self.tools] - self.param_names = { - tool["name"]: [name for name, _ in tool["parameters"]] - for tool in self.tools - } - self.param_types = { - tool["name"]: {name: _schema_type_name(param_type) for name, param_type in tool["parameters"]} - for tool in self.tools - } - self.ids = tokenizer.json_token_ids - self.quote_id = self.ids[JSON_QUOTE] - self.tool_name_trie = _build_trie( - (i, tokenizer.encode_json_string_content(name)) - for i, name in enumerate(self.tool_names) - ) - self.state = "start_array" - self.current_tool = None - self.current_tool_index = None - self.current_param_start = 0 - self.current_param_name = None - self.current_param_type = "any" - self._name_trie_node = None - self._selected_name_index = None - self.value_parser = None - self.invalid = False - - def _remaining_param_trie(self): - if self.current_tool is None: - return {} - params = self.param_names.get(self.current_tool, []) - sequences = [] - for i in range(self.current_param_start, len(params)): - sequences.append((i, self.tokenizer.encode_json_string_content(params[i]))) - return _build_trie(sequences) - - def _set_name_state(self, trie, next_state): - self.state = next_state - self._name_trie_node = trie - self._selected_name_index = None - - def allowed_next_ids(self, training=False): - if self.invalid: - return None - if self.value_parser is not None: - return self.value_parser.allowed_next_ids(training=training) - - if self.state == "done": - return np.array([EOS_ID], dtype=np.int32) - if self.state == "start_array": - return np.array([self.ids[JSON_LBRACK]], dtype=np.int32) - if self.state == "after_array_start": - return np.array([self.ids[JSON_RBRACK], self.ids[JSON_LBRACE]], dtype=np.int32) - if self.state == "after_call": - return np.array([self.ids[JSON_COMMA], self.ids[JSON_RBRACK]], dtype=np.int32) - if self.state == "expect_key_name": - return np.array([self.ids[JSON_KEY_NAME]], dtype=np.int32) - if self.state == "expect_name_colon": - return np.array([self.ids[JSON_COLON]], dtype=np.int32) - if self.state == "expect_tool_name_quote": - return np.array([self.quote_id], dtype=np.int32) - if self.state == "tool_name_content": - allowed = [token_id for token_id in self._name_trie_node.keys() if token_id != "_end"] - if "_end" in self._name_trie_node: - allowed.append(self.quote_id) - return np.array(sorted(set(allowed)), dtype=np.int32) - if self.state == "expect_tool_name_comma": - return np.array([self.ids[JSON_COMMA]], dtype=np.int32) - if self.state == "expect_key_arguments": - return np.array([self.ids[JSON_KEY_ARGUMENTS]], dtype=np.int32) - if self.state == "expect_arguments_colon": - return np.array([self.ids[JSON_COLON]], dtype=np.int32) - if self.state == "expect_arguments_object": - return np.array([self.ids[JSON_LBRACE]], dtype=np.int32) - if self.state == "arg_name_or_end": - params = self.param_names.get(self.current_tool, []) - if self.current_param_start >= len(params): - return np.array([self.ids[JSON_RBRACE]], dtype=np.int32) - return np.array([self.ids[JSON_RBRACE], self.quote_id], dtype=np.int32) - if self.state == "arg_name_content": - allowed = [token_id for token_id in self._name_trie_node.keys() if token_id != "_end"] - if "_end" in self._name_trie_node: - allowed.append(self.quote_id) - return np.array(sorted(set(allowed)), dtype=np.int32) - if self.state == "expect_arg_colon": - return np.array([self.ids[JSON_COLON]], dtype=np.int32) - if self.state == "after_arg_value": - has_more = self.current_param_start < len(self.param_names.get(self.current_tool, [])) - if has_more: - return np.array([self.ids[JSON_COMMA], self.ids[JSON_RBRACE]], dtype=np.int32) - return np.array([self.ids[JSON_RBRACE]], dtype=np.int32) - if self.state == "expect_call_end": - return np.array([self.ids[JSON_RBRACE]], dtype=np.int32) - return None - - def step(self, token_id): - token_id = int(token_id) - if self.invalid: - return - if self.value_parser is not None: - self.value_parser.step(token_id) - if self.value_parser.complete: - self.value_parser = None - self.state = "after_arg_value" - return - - if self.state == "start_array": - self.state = "after_array_start" - return - if self.state == "after_array_start": - if token_id == self.ids[JSON_RBRACK]: - self.state = "done" - else: - self.state = "expect_key_name" - return - if self.state == "after_call": - if token_id == self.ids[JSON_COMMA]: - self.state = "expect_key_name" - else: - self.state = "done" - return - if self.state == "expect_key_name": - self.state = "expect_name_colon" - return - if self.state == "expect_name_colon": - self.state = "expect_tool_name_quote" - return - if self.state == "expect_tool_name_quote": - self._set_name_state(self.tool_name_trie, "tool_name_content") - return - if self.state == "tool_name_content": - if token_id == self.quote_id: - if "_end" not in self._name_trie_node: - self.invalid = True - return - self.current_tool_index = self._name_trie_node["_end"] - self.current_tool = self.tool_names[self.current_tool_index] - self.current_param_start = 0 - self.state = "expect_tool_name_comma" - else: - if token_id not in self._name_trie_node: - self.invalid = True - return - self._name_trie_node = self._name_trie_node[token_id] - return - if self.state == "expect_tool_name_comma": - self.state = "expect_key_arguments" - return - if self.state == "expect_key_arguments": - self.state = "expect_arguments_colon" - return - if self.state == "expect_arguments_colon": - self.state = "expect_arguments_object" - return - if self.state == "expect_arguments_object": - self.state = "arg_name_or_end" - return - if self.state == "arg_name_or_end": - if token_id == self.ids[JSON_RBRACE]: - self.state = "expect_call_end" - else: - self._set_name_state(self._remaining_param_trie(), "arg_name_content") - return - if self.state == "arg_name_content": - if token_id == self.quote_id: - if "_end" not in self._name_trie_node: - self.invalid = True - return - param_idx = self._name_trie_node["_end"] - self.current_param_name = self.param_names[self.current_tool][param_idx] - self.current_param_type = self.param_types[self.current_tool].get(self.current_param_name, "any") - self.current_param_start = param_idx + 1 - self.state = "expect_arg_colon" - else: - if token_id not in self._name_trie_node: - self.invalid = True - return - self._name_trie_node = self._name_trie_node[token_id] - return - if self.state == "expect_arg_colon": - self.value_parser = JsonValueParser(self.tokenizer, type_hint=self.current_param_type) - return - if self.state == "after_arg_value": - if token_id == self.ids[JSON_COMMA]: - self.state = "arg_name_or_end" - else: - self.state = "expect_call_end" - return - if self.state == "expect_call_end": - self.state = "after_call" - - -def _extract_tools_tokens(tgt_in, loss_mask): - supervised = np.flatnonzero(loss_mask > 0) - if len(supervised) == 0: - return [] - tool_stop = int(supervised[0]) + 1 - return [int(token_id) for token_id in tgt_in[2:tool_stop] if int(token_id) != 0] - - -def build_cfg_training_constraints(tokenizer, tgt_in_batch, tgt_out_batch, loss_mask_batch): - batch_size, seq_len = tgt_out_batch.shape - batch_allowed = [[None] * seq_len for _ in range(batch_size)] - max_allowed = 0 - - for i in range(batch_size): - tools_tokens = _extract_tools_tokens(tgt_in_batch[i], loss_mask_batch[i]) - tools_text = tokenizer.decode_structured(tools_tokens) - cfg = ToolCallCFG(tokenizer, tools_text) - active_positions = np.flatnonzero(loss_mask_batch[i] > 0) - for pos in active_positions: - if cfg.invalid: - break - allowed = cfg.allowed_next_ids(training=True) - if allowed is not None and len(allowed) > 0: - allowed = np.asarray(allowed, dtype=np.int32) - batch_allowed[i][int(pos)] = allowed - max_allowed = max(max_allowed, len(allowed)) - cfg.step(int(tgt_out_batch[i, pos])) - - if max_allowed == 0: - return ( - np.full((batch_size, seq_len, 1), -1, dtype=np.int32), - np.zeros((batch_size, seq_len), dtype=np.int32), - ) - - allowed_ids = np.full((batch_size, seq_len, max_allowed), -1, dtype=np.int32) - allowed_counts = np.zeros((batch_size, seq_len), dtype=np.int32) - for i in range(batch_size): - for pos, allowed in enumerate(batch_allowed[i]): - if allowed is None: - continue - count = len(allowed) - allowed_ids[i, pos, :count] = allowed - allowed_counts[i, pos] = count - return allowed_ids, allowed_counts diff --git a/src/train.py b/src/train.py index 3d6eea8..20f089b 100644 --- a/src/train.py +++ b/src/train.py @@ -256,11 +256,9 @@ def train(args): args.max_dec_len, max_eval_samples=getattr(args, "max_eval_samples", None), ) - ckpt_path = "(skipped)" - if not getattr(args, "no_checkpoints", False): - ckpt_name = f"needle_stage2_{config.num_encoder_layers}_{config.d_model}_{global_step}.pkl" - ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) - save_checkpoint(ckpt_path, eval_params, config, extra={"stage": "stage2"}) + ckpt_name = f"needle_stage2_{config.num_encoder_layers}_{config.d_model}_{global_step}.pkl" + ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) + save_checkpoint(ckpt_path, eval_params, config, extra={"stage": "stage2"}) print(f"\n Epoch {epoch + 1}/{args.epochs}") print(f" Train loss {sum(losses) / max(len(losses), 1):.4f}") From c42aec31cec7dc2b965dc68ae5c94c5ed16c93bf Mon Sep 17 00:00:00 2001 From: Karen Mosoyan Date: Thu, 12 Mar 2026 06:33:19 +0000 Subject: [PATCH 6/6] got alright results --- src/cli.py | 8 ++ src/data.py | 123 ++++++++++++++++++++++-- src/model.py | 2 + src/tokenize_data.py | 24 ++++- src/train.py | 218 +++++++++++++++++++++++++++++++++++++++++-- 5 files changed, 354 insertions(+), 21 deletions(-) diff --git a/src/cli.py b/src/cli.py index 1bde625..aca2763 100644 --- a/src/cli.py +++ b/src/cli.py @@ -56,6 +56,12 @@ def main(): help="Number of mel frequency bins (default: 80)") p.add_argument("--max-speech-samples", type=int, default=None, help="Max voice-tool-call training samples (default: all)") + p.add_argument("--speech-replay-every", type=int, default=0, + help="Do 1 speech transcription replay step every N text steps (0=disabled)") + p.add_argument("--speech-gcs-prefix", type=str, + default="gs://cactus-dataset/speech-datav1/emilia-large") + p.add_argument("--reinit-decoder", action="store_true", + help="Reinitialise decoder weights from scratch when loading a checkpoint (keeps encoder)") p = sub.add_parser("pretrain", add_help=False) p.add_argument("--full", action="store_true") @@ -73,6 +79,8 @@ def main(): p.add_argument("--max-dec-len", type=int, default=1024) p.add_argument("--max-samples", type=int, default=None) p.add_argument("--max-speech-samples", type=int, default=None) + p.add_argument("--speech-replay-every", type=int, default=0, + help="Do 1 speech transcription replay step every N text steps (0=disabled)") p.add_argument("--warmup-ratio", type=float, default=0.05) p.add_argument("--wandb", action="store_true") p.add_argument("--dtype", type=str, default="bfloat16", choices=["float32", "bfloat16"]) diff --git a/src/data.py b/src/data.py index df97db7..72350d3 100644 --- a/src/data.py +++ b/src/data.py @@ -1358,6 +1358,70 @@ def _shard_paths(suffix): } +def prepare_audio_toolcall_val(val_data, n_mels=80, max_mel_len=1024): + """Build paired (mel, dec_in, dec_tgt, loss_mask) arrays from the tool-call val set. + + Uses the `audio` field (WAV bytes) as encoder input and the same decoder + targets as the text version, restricted to non-empty (non-trivial) examples. + Returns None if audio is unavailable for the dataset rows. + """ + import io as _io + try: + import soundfile as sf + except ImportError: + return None + + ds = _load_unified_dataset() + n_total = len(ds) + cut = int(n_total * 0.9) + val_rows = ds.select(list(range(cut, n_total))) + + kept_indices = np.array(val_data["kept_indices"]) + dec_in = np.array(val_data["dec_inputs"]) + dec_tgt = np.array(val_data["dec_targets"]) + loss_mask = np.array(val_data["loss_mask"]) + + # Filter to non-empty answers only + answer_len = loss_mask.sum(axis=1) + nonempty = np.where(answer_len >= 4)[0] + + mels = [] + idxs_kept = [] + for cache_pos in nonempty: + row_in_val = int(kept_indices[cache_pos]) + row = val_rows[row_in_val] + audio_val = row.get("audio") if isinstance(row, dict) else None + if audio_val is None: + continue + raw_bytes = audio_val.get("bytes") if isinstance(audio_val, dict) else audio_val + if not raw_bytes: + continue + try: + audio_arr, sr = sf.read(_io.BytesIO(raw_bytes), dtype="float32") + if audio_arr.ndim > 1: + audio_arr = audio_arr.mean(axis=1) + mel = compute_mel_spectrogram(audio_arr, sr=sr, n_mels=n_mels) + if mel.shape[0] > max_mel_len: + mel = mel[:max_mel_len] + elif mel.shape[0] < max_mel_len: + mel = np.pad(mel, [(0, max_mel_len - mel.shape[0]), (0, 0)]) + mels.append(mel) + idxs_kept.append(cache_pos) + except Exception: + continue + + if not mels: + return None + + idxs_kept = np.array(idxs_kept) + return { + "mels": np.stack(mels).astype(np.float32), # (N, max_mel_len, n_mels) + "dec_inputs": dec_in[idxs_kept], + "dec_targets": dec_tgt[idxs_kept], + "loss_mask": loss_mask[idxs_kept], + } + + def load_prepared_mels(mel_cache_id, mmap=False): """Load precomputed mel .npy file(s), optionally memory-mapped. @@ -1560,13 +1624,24 @@ def _list_gcs_paths(prefix): return [line.strip() for line in result.stdout.splitlines() if line.strip()] -def _download_gcs_files(uris, dest_dir, chunk_size=64): +def _download_gcs_files(uris, dest_dir, chunk_size=200, max_workers=32): + import concurrent.futures as _cf os.makedirs(dest_dir, exist_ok=True) - for start in range(0, len(uris), chunk_size): - chunk = uris[start:start + chunk_size] - if not chunk: - continue - _run_gcloud(["gcloud", "storage", "cp", *chunk, dest_dir]) + if not uris: + return + # gsutil -m cp -I only reads ~2 URIs from stdin on this environment; pass URIs as args instead. + # Use parallel invocations for throughput. + chunks = [uris[i:i + chunk_size] for i in range(0, len(uris), chunk_size)] + def _dl_chunk(chunk): + subprocess.run( + ["gsutil", "-q", "-m", + "-o", "GSUtil:parallel_process_count=1", + "-o", "GSUtil:parallel_thread_count=8", + "cp"] + list(chunk) + [dest_dir], + check=False, + ) + with _cf.ThreadPoolExecutor(max_workers=max_workers) as pool: + list(pool.map(_dl_chunk, chunks)) def _rsync_gcs_dir(src_dir, dest_dir): @@ -1786,9 +1861,6 @@ def _ensure_emilia_mels_local(mel_uris, gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, dow def prefetch_emilia_mels(mel_uris, gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, use_rsync=False): """Mirror Emilia mel .npy files into the local cache.""" _, mel_dir = _emilia_cache_dirs(gcs_prefix) - if use_rsync: - _rsync_gcs_dir(f"{gcs_prefix.rstrip('/')}/train/mels", mel_dir) - return sum(1 for name in os.listdir(mel_dir) if name.endswith(".npy")) unique_uris = [] seen = set() @@ -1797,7 +1869,38 @@ def prefetch_emilia_mels(mel_uris, gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, use_rsyn if uri not in seen: seen.add(uri) unique_uris.append(uri) - _ensure_emilia_mels_local(unique_uris, gcs_prefix=gcs_prefix, download_missing=True) + + print(f"Stage 1 unique mel files selected: {len(unique_uris):,}") + + if use_rsync: + print("Stage 1 mel mirror mode: full-rsync") + print(f" target: {mel_dir}") + print(" (ignores --max-speech-samples)") + _rsync_gcs_dir(f"{gcs_prefix.rstrip('/')}/train/mels", mel_dir) + n = sum(1 for name in os.listdir(mel_dir) if name.endswith(".npy")) + print(f"Stage 1 local mel cache: {n:,} files after rsync") + return n + + print("Stage 1 mel mirror mode: subset-copy") + local_files = set(os.listdir(mel_dir)) + missing = [uri for uri in unique_uris if os.path.basename(uri) not in local_files] + n_reused = len(unique_uris) - len(missing) + print(f"Stage 1 local mel cache reused: {n_reused:,}") + print(f"Stage 1 local mel files to copy: {len(missing):,}") + + if missing: + import time as _time + t0 = _time.perf_counter() + _download_gcs_files(missing, mel_dir) + dt = _time.perf_counter() - t0 + n_copied = sum( + 1 for uri in missing if os.path.exists(os.path.join(mel_dir, os.path.basename(uri))) + ) + mb = n_copied * 150 / 1024 + print(f"Stage 1 local mel files copied: {n_copied:,} ({mb:.0f} MB in {dt:.1f}s, {mb/max(dt,0.01):.1f} MB/s)") + else: + print("Stage 1 local mel files copied: 0 (all already present)") + return len(unique_uris) diff --git a/src/model.py b/src/model.py index bc0c468..ef642c5 100644 --- a/src/model.py +++ b/src/model.py @@ -391,6 +391,8 @@ def _run_decoder(self, encoder_out, tgt, tgt_mask=None, ffn_mask=None, determini def _slot_diversity(self, encoder_out): s = encoder_out.astype(jnp.float32) + # Normalize to unit sphere so diversity is scale-invariant + s = s / (jnp.linalg.norm(s, axis=-1, keepdims=True) + 1e-8) gram = jnp.matmul(s, s.transpose(0, 2, 1)) diag_sq = jnp.sum(jnp.diagonal(gram, axis1=1, axis2=2) ** 2) return (jnp.sum(gram ** 2) - diag_sq) / s.shape[0] diff --git a/src/tokenize_data.py b/src/tokenize_data.py index 5484bca..cfef6e3 100644 --- a/src/tokenize_data.py +++ b/src/tokenize_data.py @@ -47,11 +47,22 @@ def _clear_gcs_tokenizer(): def _clear_local_caches(): - """Remove local cache and tokenizer directories.""" - for d in [CACHE_DIR, TOKENIZER_DIR]: - if os.path.exists(d): - print(f"Removing {d}/ ...") - shutil.rmtree(d) + """Remove local tokenizer and text caches, but preserve Emilia mel mirrors.""" + if os.path.exists(TOKENIZER_DIR): + print(f"Removing {TOKENIZER_DIR}/ ...") + shutil.rmtree(TOKENIZER_DIR) + + if os.path.exists(CACHE_DIR): + for item in sorted(os.listdir(CACHE_DIR)): + path = os.path.join(CACHE_DIR, item) + if item.startswith("emilia_"): + print(f"Preserving {path}/ (Emilia mel mirror)") + continue + if os.path.isdir(path): + print(f"Removing {path}/ ...") + shutil.rmtree(path) + else: + os.remove(path) def tokenize(args): @@ -84,6 +95,7 @@ def tokenize(args): speech_val_ratio = getattr(args, "speech_val_ratio", 0.01) speech_max_samples = getattr(args, "max_speech_samples", None) or args.max_samples stage1_mel_uris = [] + stage1_counts = {} for split in ("train", "val"): rows = load_emilia_speech_metadata( split, @@ -92,6 +104,7 @@ def tokenize(args): val_ratio=speech_val_ratio, seed=getattr(args, "split_seed", 42), ) + stage1_counts[split] = len(rows) prepared = prepare_transcription_pairs(rows, tokenizer, args.max_enc_len, args.max_dec_len) cache_id = save_prepared_transcription_data( split, @@ -104,6 +117,7 @@ def tokenize(args): ) stage1_mel_uris.extend(prepared["mel_uris"].tolist()) print(f"Cached {len(rows):,} Emilia {split} examples ({cache_id})") + print(f"Stage 1 Emilia rows selected: {stage1_counts['train']:,} train / {stage1_counts['val']:,} val") print("\n=== Mirroring Emilia mel files for Stage 1 ===") use_rsync = speech_max_samples is None diff --git a/src/train.py b/src/train.py index 20f089b..84dd2db 100644 --- a/src/train.py +++ b/src/train.py @@ -9,8 +9,11 @@ from flax import jax_utils from tqdm import tqdm -from .data import PrefetchIterator, count_batches, get_batches, get_tokenizer, load_prepared_data -from .model import EncoderDecoderTransformer, make_causal_mask, make_padding_mask +from .data import ( + PrefetchIterator, count_batches, get_batches, get_tokenizer, load_prepared_data, + prepare_audio_toolcall_val, +) +from .model import EncoderDecoderTransformer, make_causal_mask, make_mel_padding_mask, make_padding_mask from .train_utils import ( count_params, create_config_from_args, @@ -67,6 +70,46 @@ def _make_p_train_step(): return jax.pmap(_train_step, axis_name="batch", donate_argnums=(0, 1)) +def _speech_replay_loss_fn(state, params, mel, tgt_in, tgt_out, causal_mask, rng, loss_mask): + """Speech transcription loss for decoder-only replay (encoder frozen via zero_non_decoder_grads).""" + pad_id = 0 + src_mask = make_mel_padding_mask(mel) + tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) + logits, _ = state.apply_fn( + {"params": quantize_params(params, group_size=_GROUP_SIZE)}, + mel, + tgt_in, + src_mask=src_mask, + tgt_mask=tgt_mask, + deterministic=False, + method="forward_speech_masked", + rngs={"dropout": rng}, + ) + logits_f32 = logits.astype(jnp.float32) + token_loss = optax.softmax_cross_entropy_with_integer_labels(logits_f32, tgt_out) + ce_loss = jnp.sum(token_loss * loss_mask) / jnp.maximum(jnp.sum(loss_mask), 1.0) + z_loss = 1e-4 * jnp.mean(jax.nn.logsumexp(logits_f32, axis=-1) ** 2) + return ce_loss + z_loss + + +def _speech_replay_step(state, ema_params, mel, tgt_in, tgt_out, causal_mask, rng, loss_mask): + ema_decay = 0.999 + loss, grads = jax.value_and_grad( + lambda p: _speech_replay_loss_fn(state, p, mel, tgt_in, tgt_out, causal_mask, rng, loss_mask) + )(state.params) + grads = jax.lax.pmean(grads, axis_name="batch") + grads = zero_non_decoder_grads(grads) + loss = jax.lax.pmean(loss, axis_name="batch") + grad_norm = optax.global_norm(grads) + state = state.apply_gradients(grads=grads) + ema_params = jax.tree.map(lambda e, p: ema_decay * e + (1 - ema_decay) * p, ema_params, state.params) + return state, ema_params, loss, grad_norm + + +def _make_p_speech_replay_step(): + return jax.pmap(_speech_replay_step, axis_name="batch", donate_argnums=(0, 1)) + + def _make_val_loss_fn(apply_fn): @jax.jit def val_loss_batch(params, src, tgt_in, tgt_out, causal_mask, loss_mask): @@ -97,6 +140,70 @@ def _evaluate_val_ppl(val_loss_fn, params, val_enc, val_dec_in, val_dec_tgt, val return float(math.exp(min(total_loss / max(total_toks, 1.0), 20.0))) +def _filter_nonempty(val_enc, val_dec_in, val_dec_tgt, val_loss_mask, min_answer_tokens=4): + """Return val arrays filtered to examples with at least min_answer_tokens answer tokens.""" + answer_len = np.array(val_loss_mask).sum(axis=1) + mask = answer_len >= min_answer_tokens + idx = np.where(mask)[0] + return (val_enc[idx], val_dec_in[idx], val_dec_tgt[idx], val_loss_mask[idx]) + + +def _make_speech_val_loss_fn(apply_fn): + """Val loss using frozen speech encoder path (forward_speech_masked).""" + @jax.jit + def speech_val_loss_batch(params, mel, tgt_in, tgt_out, causal_mask, loss_mask): + pad_id = 0 + src_mask = make_mel_padding_mask(mel) + tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) + logits, _ = apply_fn( + {"params": params}, + mel, + tgt_in, + src_mask=src_mask, + tgt_mask=tgt_mask, + deterministic=True, + method="forward_speech_masked", + ) + loss = optax.softmax_cross_entropy_with_integer_labels(logits.astype(jnp.float32), tgt_out) + return jnp.sum(loss * loss_mask), jnp.sum(loss_mask) + + return speech_val_loss_batch + + +def _evaluate_audio_toolcall_ppl(speech_val_loss_fn, params, audio_val_data, + batch_size, max_dec_len, max_eval_samples=None): + """Evaluate tool-call PPL using audio (mel) encoder input — the correct speech transfer metric. + + Uses speech query audio (from dataset `audio` field) → memory slots → tool call answer. + Same decoder targets as the text val path, so directly comparable. + """ + val_causal = make_causal_mask(max_dec_len) + mels = audio_val_data["mels"] + dec_in = audio_val_data["dec_inputs"] + dec_tgt = audio_val_data["dec_targets"] + loss_mask = audio_val_data["loss_mask"] + + total_loss = 0.0 + total_toks = 0.0 + seen = 0 + n = len(mels) + for start in range(0, n, batch_size): + if max_eval_samples is not None and seen >= max_eval_samples: + break + end = min(start + batch_size, n) + mel_b = mels[start:end] + di_b = dec_in[start:end] + dt_b = dec_tgt[start:end] + lm_b = loss_mask[start:end] + vl, vt = speech_val_loss_fn(params, mel_b, di_b, dt_b, val_causal, lm_b) + total_loss += float(vl) + total_toks += float(vt) + seen += len(mel_b) + if total_toks == 0: + return None + return float(math.exp(min(total_loss / total_toks, 20.0))) + + def train(args): global _GROUP_SIZE _GROUP_SIZE = getattr(args, "group_size", 32) @@ -123,13 +230,27 @@ def train(args): val_dec_in = val_data["dec_inputs"] val_dec_tgt = val_data["dec_targets"] val_loss_mask = val_data["loss_mask"] + nonempty_val_enc, nonempty_val_dec_in, nonempty_val_dec_tgt, nonempty_val_loss_mask = \ + _filter_nonempty(val_enc, val_dec_in, val_dec_tgt, val_loss_mask) if args.max_samples is not None: keep = min(args.max_samples, len(enc_inputs)) enc_inputs = enc_inputs[:keep] dec_inputs = dec_inputs[:keep] dec_targets = dec_targets[:keep] train_loss_mask = train_loss_mask[:keep] - print(f" {len(enc_inputs):,} train / {len(val_enc):,} val tool-call pairs") + print(f" {len(enc_inputs):,} train / {len(val_enc):,} val tool-call pairs " + f"({len(nonempty_val_enc):,} non-empty val)") + + # Load audio+tool-call val data for speech transfer eval + # (speech query audio → same tool-call targets as text val, measures audio→tool-call PPL) + n_mels = getattr(args, "n_mels", 80) + max_mel_len = getattr(args, "max_mel_len", 1024) + print(f" Preparing audio+tool-call val data for speech transfer eval...") + audio_toolcall_val = prepare_audio_toolcall_val(val_data, n_mels=n_mels, max_mel_len=max_mel_len) + if audio_toolcall_val is not None: + print(f" {len(audio_toolcall_val['mels']):,} audio+tool-call val pairs ready") + else: + print(f" Audio+tool-call val data unavailable") print(f"\n[3/3] Building model...") resume_checkpoint = getattr(args, "checkpoint", None) @@ -153,10 +274,21 @@ def train(args): muon_lr = getattr(args, "muon_lr", 0.02) * math.sqrt(num_devices) state = create_train_state(init_rng, config, scaled_lr, muon_lr, total_steps, warmup_steps) if ckpt_params is not None: - state = state.replace(params=ckpt_params) + if getattr(args, "reinit_decoder", False): + # Keep encoder/embedding/mel_proj from checkpoint; reinitialise decoder from scratch. + fresh_params = state.params + merged = {k: (ckpt_params[k] if k != "decoder" else fresh_params[k]) + for k in ckpt_params} + state = state.replace(params=merged) + print(f" decoder reinitialised from scratch (encoder weights kept)") + else: + state = state.replace(params=ckpt_params) val_loss_fn = _make_val_loss_fn(state.apply_fn) + speech_val_loss_fn = _make_speech_val_loss_fn(state.apply_fn) if audio_toolcall_val is not None else None p_train_step = _make_p_train_step() + speech_replay_every = getattr(args, "speech_replay_every", 0) + p_speech_replay_step = _make_p_speech_replay_step() if speech_replay_every > 0 else None ema_params = jax.tree.map(jnp.copy, state.params) state = jax_utils.replicate(state) @@ -181,8 +313,14 @@ def train(args): global_step = 0 last_val_ppl = None + last_nonempty_ppl = None + last_speech_ppl = None + best_nonempty_ppl = float("inf") eval_every = getattr(args, "eval_every", 1000) + # Speech replay (optional, disabled by default) + speech_replay_iter = None + for epoch in range(args.epochs): losses = [] batch_iter = PrefetchIterator( @@ -210,6 +348,24 @@ def train(args): losses.append(loss_val) global_step += 1 + # Speech replay step (decoder-only, prevents catastrophic forgetting) + if (p_speech_replay_step is not None and speech_replay_every > 0 + and global_step % speech_replay_every == 0): + try: + speech_batch = next(speech_replay_iter) + mel_b, _, sp_tgt_in, sp_tgt_out, sp_lm = speech_batch + mel_b = shard_batch(mel_b, num_devices) + sp_tgt_in_b = shard_batch(sp_tgt_in, num_devices) + sp_tgt_out_b = shard_batch(sp_tgt_out, num_devices) + sp_lm_b = shard_batch(sp_lm, num_devices) + rng, sp_rng = jax.random.split(rng) + sp_rngs = jax.random.split(sp_rng, num_devices) + state, ema_params, _, _ = p_speech_replay_step( + state, ema_params, mel_b, sp_tgt_in_b, sp_tgt_out_b, causal_mask, sp_rngs, sp_lm_b + ) + except StopIteration: + pass + if global_step % eval_every == 0 or global_step == total_steps: eval_params = jax_utils.unreplicate(ema_params) last_val_ppl = _evaluate_val_ppl( @@ -223,10 +379,36 @@ def train(args): args.max_dec_len, max_eval_samples=getattr(args, "max_eval_samples", None), ) + last_nonempty_ppl = _evaluate_val_ppl( + val_loss_fn, + eval_params, + nonempty_val_enc, + nonempty_val_dec_in, + nonempty_val_dec_tgt, + nonempty_val_loss_mask, + args.batch_size, + args.max_dec_len, + max_eval_samples=getattr(args, "max_eval_samples", None), + ) + if speech_val_loss_fn is not None: + last_speech_ppl = _evaluate_audio_toolcall_ppl( + speech_val_loss_fn, + eval_params, + audio_toolcall_val, + args.batch_size, + max_dec_len=args.max_dec_len, + max_eval_samples=getattr(args, "max_eval_samples", None), + ) + if last_nonempty_ppl < best_nonempty_ppl: + best_nonempty_ppl = last_nonempty_ppl + best_ckpt = os.path.join(args.checkpoint_dir, f"needle_stage2_{config.num_encoder_layers}_{config.d_model}_best.pkl") + save_checkpoint(best_ckpt, eval_params, config, extra={"stage": "stage2", "step": global_step, "val_ppl": last_val_ppl, "nonempty_val_ppl": best_nonempty_ppl, "speech_ppl": last_speech_ppl}) pbar.set_postfix( loss=f"{loss_val:.4f}", ppl=f"{last_val_ppl:.2f}" if last_val_ppl is not None else "?", + ne_ppl=f"{last_nonempty_ppl:.2f}" if last_nonempty_ppl is not None else "?", + aud_ppl=f"{last_speech_ppl:.2f}" if last_speech_ppl is not None else "?", ) if use_wandb: @@ -238,7 +420,8 @@ def train(args): "train/muon_lr": float(muon_schedule(global_step)), "train/tokens_per_sec": effective_batch_size * (args.max_enc_len + args.max_dec_len) / max(dt, 1e-6), "train/step": global_step, - **({"val/ppl": last_val_ppl} if last_val_ppl is not None and (global_step % eval_every == 0 or global_step == total_steps) else {}), + **({"val/ppl": last_val_ppl, "val/nonempty_ppl": last_nonempty_ppl} + if last_val_ppl is not None and (global_step % eval_every == 0 or global_step == total_steps) else {}), } ) @@ -256,13 +439,36 @@ def train(args): args.max_dec_len, max_eval_samples=getattr(args, "max_eval_samples", None), ) + nonempty_ppl = _evaluate_val_ppl( + val_loss_fn, + eval_params, + nonempty_val_enc, + nonempty_val_dec_in, + nonempty_val_dec_tgt, + nonempty_val_loss_mask, + args.batch_size, + args.max_dec_len, + max_eval_samples=getattr(args, "max_eval_samples", None), + ) + speech_ppl = None + if speech_val_loss_fn is not None: + speech_ppl = _evaluate_audio_toolcall_ppl( + speech_val_loss_fn, + eval_params, + audio_toolcall_val, + args.batch_size, + max_dec_len=args.max_dec_len, + max_eval_samples=getattr(args, "max_eval_samples", None), + ) ckpt_name = f"needle_stage2_{config.num_encoder_layers}_{config.d_model}_{global_step}.pkl" ckpt_path = os.path.join(args.checkpoint_dir, ckpt_name) save_checkpoint(ckpt_path, eval_params, config, extra={"stage": "stage2"}) print(f"\n Epoch {epoch + 1}/{args.epochs}") print(f" Train loss {sum(losses) / max(len(losses), 1):.4f}") - print(f" Val ppl {val_ppl:.2f}") + print(f" Val ppl {val_ppl:.2f} (all) {nonempty_ppl:.2f} (non-empty only)") + if speech_ppl is not None: + print(f" Audio val ppl {speech_ppl:.2f} (speech query → tool call)") print(f" Checkpoint {ckpt_path}\n") if use_wandb: