diff --git a/setup b/setup index 26b3307..6511ee7 100755 --- a/setup +++ b/setup @@ -49,6 +49,11 @@ fi source "$VENV_DIR/bin/activate" +export NEEDLE_CACHE_DIR="${NEEDLE_CACHE_DIR:-/dev/shm/needle_cache}" +export NEEDLE_TOKENIZER_DIR="${NEEDLE_TOKENIZER_DIR:-/dev/shm/needle_tokenizer}" +export NEEDLE_LOCAL_UNIFIED_DIR="${NEEDLE_LOCAL_UNIFIED_DIR:-/dev/shm/needle_tool_calls_unified}" +mkdir -p "$NEEDLE_CACHE_DIR" "$NEEDLE_TOKENIZER_DIR" "$NEEDLE_LOCAL_UNIFIED_DIR" + echo "Installing dependencies..." pip install --upgrade pip -q pip install -e . -q diff --git a/src/cli.py b/src/cli.py index 15a909e..aca2763 100644 --- a/src/cli.py +++ b/src/cli.py @@ -56,6 +56,55 @@ 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") + p.add_argument("--checkpoint", type=str, default=None) + p.add_argument("--epochs", type=int, default=3) + p.add_argument("--batch-size", type=int, default=32) + p.add_argument("--lr", type=float, default=3e-4) + p.add_argument("--muon-lr", type=float, default=0.02) + p.add_argument("--d-model", type=int, default=512) + p.add_argument("--num-heads", type=int, default=16) + p.add_argument("--num-kv-heads", type=int, default=8) + p.add_argument("--num-layers", type=int, default=4) + p.add_argument("--num-dec-layers", type=int, default=4) + p.add_argument("--max-enc-len", type=int, default=256) + 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"]) + p.add_argument("--checkpoint-dir", type=str, default="checkpoints") + 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("--group-size", type=int, default=32) + 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("--dropout", type=float, default=0.0, + help="Dropout rate for residual connections (default: 0.1)") + 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("--speech-gcs-prefix", type=str, default="gs://cactus-dataset/speech-datav1/emilia-large") + p.add_argument("--speech-val-ratio", type=float, default=0.01) + p.add_argument("--toucan-config", type=str, default="Kimi-K2") + p.add_argument("--toucan-max-samples", type=int, default=None) + p.add_argument("--tool-contrastive-weight", type=float, default=1.0, + 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 = sub.add_parser("tokenize", add_help=False) p.add_argument("--max-samples", type=int, default=None, @@ -72,6 +121,20 @@ 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.add_argument("--speech-gcs-prefix", type=str, default="gs://cactus-dataset/speech-datav1/emilia-large") + p.add_argument("--speech-val-ratio", type=float, default=0.01) + p.add_argument("--max-speech-samples", type=int, default=None, + help="Optional cap for Emilia Stage 1 preprocessing") + p.add_argument("--toucan-config", type=str, default="Kimi-K2", + 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("--overwrite-gcs-tokenizer", action="store_true", + help="Overwrite the shared GCS tokenizer after retraining locally") p = sub.add_parser("run", add_help=False) p.add_argument("--checkpoint", type=str, required=True) @@ -140,6 +203,16 @@ def main(): if args.command == "tokenize": from .tokenize_data import tokenize tokenize(args) + elif args.command == "pretrain": + if getattr(args, "full", False): + args.d_model = 1536 + args.num_heads = 24 + args.num_kv_heads = 8 + args.num_layers = 12 + args.num_dec_layers = 4 + args.num_memory_slots = 128 + from .pretrain import pretrain + pretrain(args) elif args.command == "train": if getattr(args, "full", False): args.d_model = 1536 diff --git a/src/data.py b/src/data.py index b916035..72350d3 100644 --- a/src/data.py +++ b/src/data.py @@ -1,8 +1,10 @@ +import csv import hashlib import json as _json import multiprocessing as mp import os import queue +import subprocess import threading os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "1" os.environ["HF_HUB_DISABLE_TELEMETRY"] = "1" @@ -12,6 +14,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 @@ -19,16 +22,21 @@ _PROJECT_ROOT = os.path.dirname(os.path.dirname(__file__)) _DISK_CACHE_DIR = os.path.join(_PROJECT_ROOT, ".data_cache") -TOKENIZER_DIR = os.path.join(_PROJECT_ROOT, "tokenizer") -TOKENIZER_PREFIX = os.path.join(TOKENIZER_DIR, "needle") -LOCAL_UNIFIED_DIR = os.path.join(_PROJECT_ROOT, "data", "tool_calls_unified") +_DEFAULT_TOKENIZER_DIR = os.path.join(_PROJECT_ROOT, "tokenizer") +_DEFAULT_LOCAL_UNIFIED_DIR = os.path.join(_PROJECT_ROOT, "data", "tool_calls_unified") GCS_DATASET_PATH = "gs://cactus-dataset/tool_calls" +EMILIA_SPEECH_GCS_PREFIX = "gs://cactus-dataset/speech-datav1/emilia-large" _MIN_SHM_BYTES = 200 * 1024**3 def _pick_cache_dir(): """Use /dev/shm (tmpfs/RAM) when available and large enough, else disk.""" + env_cache_dir = os.environ.get("NEEDLE_CACHE_DIR") + if env_cache_dir: + os.makedirs(env_cache_dir, exist_ok=True) + return env_cache_dir + shm = "/dev/shm" if os.path.isdir(shm): try: @@ -45,6 +53,9 @@ def _pick_cache_dir(): CACHE_DIR = _pick_cache_dir() +TOKENIZER_DIR = os.environ.get("NEEDLE_TOKENIZER_DIR", _DEFAULT_TOKENIZER_DIR) +TOKENIZER_PREFIX = os.path.join(TOKENIZER_DIR, "needle") +LOCAL_UNIFIED_DIR = os.environ.get("NEEDLE_LOCAL_UNIFIED_DIR", _DEFAULT_LOCAL_UNIFIED_DIR) PAD_ID = 0 EOS_ID = 1 @@ -53,6 +64,43 @@ def _pick_cache_dir(): 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 @@ -102,12 +150,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): @@ -136,10 +370,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 = [] @@ -186,10 +469,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, @@ -200,7 +490,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(), @@ -303,28 +593,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 @@ -352,8 +636,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, toucan_cache_path=None): """Save metadata JSON for a split, upload to GCS.""" os.makedirs(CACHE_DIR, exist_ok=True) meta = { @@ -365,6 +671,10 @@ 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, + "toucan_cache_path": toucan_cache_path, } meta_path = os.path.join(CACHE_DIR, f"{split}_metadata.json") with open(meta_path, "w") as f: @@ -461,19 +771,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) @@ -632,26 +940,29 @@ 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): +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() _GCS_SHARD_ROWS = 5000 @@ -1019,6 +1330,10 @@ 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) + 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"] @@ -1036,6 +1351,74 @@ 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), + "toucan_cache_path": meta.get("toucan_cache_path"), + } + + +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], } @@ -1214,3 +1597,358 @@ def close(self): def count_batches(n_samples, batch_size): """Return the number of full batches for a dataset of n_samples.""" return n_samples // batch_size + + +def _gcs_slug(path): + return hashlib.md5(path.encode()).hexdigest()[:12] + + +def _emilia_cache_dirs(gcs_prefix): + root = os.path.join(CACHE_DIR, f"emilia_{_gcs_slug(gcs_prefix)}") + meta_dir = os.path.join(root, "metadata") + mel_dir = os.path.join(root, "mels") + os.makedirs(meta_dir, exist_ok=True) + os.makedirs(mel_dir, exist_ok=True) + return meta_dir, mel_dir + + +def _run_gcloud(args): + result = subprocess.run(args, capture_output=True, text=True) + if result.returncode != 0: + raise RuntimeError(result.stderr.strip() or result.stdout.strip() or "gcloud command failed") + return result + + +def _list_gcs_paths(prefix): + result = _run_gcloud(["gcloud", "storage", "ls", prefix]) + return [line.strip() for line in result.stdout.splitlines() if line.strip()] + + +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) + 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): + os.makedirs(dest_dir, exist_ok=True) + _run_gcloud(["gcloud", "storage", "rsync", src_dir, dest_dir, "--recursive"]) + + +def _ensure_emilia_metadata_local(gcs_prefix): + meta_dir, _ = _emilia_cache_dirs(gcs_prefix) + remote_prefix = f"{gcs_prefix.rstrip('/')}/train/metadata/" + remote_paths = _list_gcs_paths(remote_prefix) + missing = [] + for uri in remote_paths: + local_path = os.path.join(meta_dir, os.path.basename(uri)) + if not os.path.exists(local_path): + missing.append(uri) + if missing: + _download_gcs_files(missing, meta_dir, chunk_size=32) + return [os.path.join(meta_dir, os.path.basename(uri)) for uri in remote_paths] + + +def load_emilia_speech_metadata(split="train", gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, + max_samples=None, val_ratio=0.01, seed=42): + """Load Emilia speech metadata rows and create a deterministic train/val split.""" + if split not in ("train", "val", "validation", "test"): + raise ValueError(f"Unsupported split: {split}") + + local_csvs = sorted(_ensure_emilia_metadata_local(gcs_prefix)) + rows = [] + for csv_path in local_csvs: + with open(csv_path, newline="") as f: + reader = csv.DictReader(f) + for row in reader: + transcript = (row.get("transcript") or "").strip() + mel_uri = (row.get("mel_gcs_uri") or "").strip() + if not transcript or not mel_uri.endswith(".npy"): + continue + rows.append(row) + + rng = np.random.default_rng(seed) + order = rng.permutation(len(rows)) + cut = int(len(rows) * (1.0 - val_ratio)) + if split == "train": + selected = order[:cut] + else: + selected = order[cut:] + + if max_samples is not None: + selected = selected[:min(max_samples, len(selected))] + + return [rows[int(i)] for i in selected] + + +def prepare_transcription_pairs(rows, tokenizer, max_enc_len=256, max_dec_len=1024): + """Prepare audio->transcript decoder targets plus transcript text encodings.""" + pad_id = tokenizer.pad_token_id + eos_id = tokenizer.eos_token_id + transcribe_id = tokenizer.transcribe_token_id + + n = len(rows) + text_inputs = np.full((n, max_enc_len), pad_id, dtype=np.int32) + dec_inputs = np.full((n, max_dec_len), pad_id, dtype=np.int32) + dec_targets = np.full((n, max_dec_len), pad_id, dtype=np.int32) + loss_mask = np.zeros((n, max_dec_len), dtype=np.float32) + mel_uris = [] + + for i, row in enumerate(rows): + transcript = row["transcript"].strip() + text_tokens = tokenizer.encode(transcript)[:max_enc_len] + if text_tokens: + text_inputs[i, :len(text_tokens)] = text_tokens + + dec_tokens = tokenizer.encode(transcript)[:max(0, max_dec_len - 2)] + dec_inputs[i, 0] = eos_id + dec_inputs[i, 1] = transcribe_id + if dec_tokens: + dec_inputs[i, 2:2 + len(dec_tokens)] = dec_tokens + + dec_targets[i, 0] = transcribe_id + if dec_tokens: + dec_targets[i, 1:1 + len(dec_tokens)] = dec_tokens + eos_pos = 1 + len(dec_tokens) + if eos_pos < max_dec_len: + dec_targets[i, eos_pos] = eos_id + loss_mask[i, 1:eos_pos + 1] = 1.0 + + mel_uris.append(row["mel_gcs_uri"]) + + return { + "text_inputs": text_inputs, + "dec_inputs": dec_inputs, + "dec_targets": dec_targets, + "loss_mask": loss_mask, + "mel_uris": np.array(mel_uris, dtype=object), + } + + +def _transcription_cache_key(split, n_samples, max_enc_len, max_dec_len, gcs_prefix, val_ratio): + tok_hash = _tokenizer_hash() + key = ( + f"transcribe_{split}_{_gcs_slug(gcs_prefix)}_{tok_hash}_" + f"{n_samples}_{max_enc_len}_{max_dec_len}_{val_ratio:.6f}" + ) + return hashlib.md5(key.encode()).hexdigest()[:12] + + +def save_prepared_transcription_data(split, prepared, max_enc_len, max_dec_len, + gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, val_ratio=0.01, + split_max_samples=None): + """Persist tokenized Emilia transcript targets for Stage 1.""" + os.makedirs(CACHE_DIR, exist_ok=True) + n_samples = len(prepared["text_inputs"]) + cache_id = _transcription_cache_key( + split, n_samples, max_enc_len, max_dec_len, gcs_prefix, val_ratio + ) + cache_path = os.path.join(CACHE_DIR, cache_id) + + np.save(cache_path + "_text_inputs.npy", prepared["text_inputs"]) + np.save(cache_path + "_dec_inputs.npy", prepared["dec_inputs"]) + np.save(cache_path + "_dec_targets.npy", prepared["dec_targets"]) + np.save(cache_path + "_loss_mask.npy", prepared["loss_mask"]) + np.save(cache_path + "_mel_uris.npy", np.asarray(prepared["mel_uris"], dtype=np.str_)) + _gcs_cache_upload( + cache_id, + ["_text_inputs.npy", "_dec_inputs.npy", "_dec_targets.npy", "_loss_mask.npy", "_mel_uris.npy"], + ) + + meta = { + "split": split, + "cache_id": cache_id, + "n_samples": n_samples, + "max_enc_len": max_enc_len, + "max_dec_len": max_dec_len, + "speech_gcs_prefix": gcs_prefix, + "speech_val_ratio": val_ratio, + "split_max_samples": split_max_samples, + } + meta_path = os.path.join(CACHE_DIR, f"stage1_{split}_metadata.json") + with open(meta_path, "w") as f: + _json.dump(meta, f) + _gcs_cache_upload(f"stage1_{split}_metadata", [".json"]) + return cache_id + + +def load_prepared_transcription_data(split, max_enc_len, max_dec_len, + gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, val_ratio=0.01, + max_samples=None, mmap=False): + """Load cached Stage 1 transcription targets if metadata matches.""" + meta_path = os.path.join(CACHE_DIR, f"stage1_{split}_metadata.json") + if not os.path.exists(meta_path): + import subprocess + subprocess.run( + ["gcloud", "storage", "cp", f"{GCS_CACHE_PATH}/stage1_{split}_metadata.json", meta_path], + capture_output=True, text=True, + ) + if not os.path.exists(meta_path): + return None + + with open(meta_path) as f: + meta = _json.load(f) + + if meta.get("speech_gcs_prefix") != gcs_prefix: + return None + if meta.get("max_enc_len") != max_enc_len or meta.get("max_dec_len") != max_dec_len: + return None + if abs(float(meta.get("speech_val_ratio", 0.01)) - float(val_ratio)) > 1e-12: + return None + + cache_id = meta["cache_id"] + cache_path = os.path.join(CACHE_DIR, cache_id) + suffixes = [ + "_text_inputs.npy", + "_dec_inputs.npy", + "_dec_targets.npy", + "_loss_mask.npy", + "_mel_uris.npy", + ] + if not all(os.path.exists(cache_path + suffix) for suffix in suffixes): + if not _gcs_cache_download(cache_id, suffixes): + return None + + mmap_mode = "r" if mmap else None + prepared = { + "text_inputs": np.load(cache_path + "_text_inputs.npy", mmap_mode=mmap_mode), + "dec_inputs": np.load(cache_path + "_dec_inputs.npy", mmap_mode=mmap_mode), + "dec_targets": np.load(cache_path + "_dec_targets.npy", mmap_mode=mmap_mode), + "loss_mask": np.load(cache_path + "_loss_mask.npy", mmap_mode=mmap_mode), + "mel_uris": np.load(cache_path + "_mel_uris.npy", mmap_mode=mmap_mode), + } + + if max_samples is not None: + keep = min(max_samples, len(prepared["text_inputs"])) + prepared = {key: np.array(value[:keep]) for key, value in prepared.items()} + + return prepared + + +def _ensure_emilia_mels_local(mel_uris, gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, download_missing=True): + _, mel_dir = _emilia_cache_dirs(gcs_prefix) + missing = [] + local_paths = [] + for uri in mel_uris: + local_path = os.path.join(mel_dir, os.path.basename(uri)) + local_paths.append(local_path) + if not os.path.exists(local_path): + missing.append(uri) + if missing and download_missing: + _download_gcs_files(missing, mel_dir, chunk_size=32) + elif missing: + raise FileNotFoundError( + f"Missing {len(missing)} local Emilia mel files under {mel_dir}. " + "Run 'needle tokenize' to mirror Stage 1 mels locally before pretraining." + ) + return local_paths + + +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) + + unique_uris = [] + seen = set() + for uri in mel_uris: + uri = str(uri) + if uri not in seen: + seen.add(uri) + unique_uris.append(uri) + + 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) + + +def _load_emilia_mel(local_path, n_mels=80, max_mel_len=1024): + mel = np.load(local_path) + if mel.ndim != 2: + raise ValueError(f"Expected 2D mel array, got shape {mel.shape} for {local_path}") + if mel.shape[0] == n_mels and mel.shape[1] != n_mels: + mel = mel.T + elif mel.shape[1] != n_mels: + raise ValueError(f"Unexpected mel shape {mel.shape} for {local_path}") + mel = mel.astype(np.float32) + 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))) + return mel + + +def get_transcription_batches(prepared, batch_size, max_mel_len=1024, n_mels=80, + gcs_prefix=EMILIA_SPEECH_GCS_PREFIX, shuffle=True, + require_local_mels=False): + """Yield batches of cached Emilia mels and transcript targets.""" + text_inputs = prepared["text_inputs"] + dec_inputs = prepared["dec_inputs"] + dec_targets = prepared["dec_targets"] + loss_mask = prepared["loss_mask"] + mel_uris = prepared["mel_uris"] + + n = len(text_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] + batch_uris = [str(mel_uris[j]) for j in idx] + local_paths = _ensure_emilia_mels_local( + batch_uris, + gcs_prefix=gcs_prefix, + download_missing=not require_local_mels, + ) + mel_batch = np.stack([ + _load_emilia_mel(path, n_mels=n_mels, max_mel_len=max_mel_len) + for path in local_paths + ]).astype(np.float32) + yield ( + mel_batch, + np.array(text_inputs[idx]), + np.array(dec_inputs[idx]), + np.array(dec_targets[idx]), + np.array(loss_mask[idx]), + ) 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/pretrain.py b/src/pretrain.py new file mode 100644 index 0000000..1c76a98 --- /dev/null +++ b/src/pretrain.py @@ -0,0 +1,442 @@ +import math +import os +import time + +import jax +import jax.numpy as jnp +import numpy as np +import optax +from flax import jax_utils +from tqdm import tqdm + +from .data import ( + EMILIA_SPEECH_GCS_PREFIX, + PrefetchIterator, + count_batches, + get_tokenizer, + get_transcription_batches, + load_emilia_speech_metadata, + load_prepared_transcription_data, + prepare_transcription_pairs, +) +from .model import make_causal_mask, make_mel_padding_mask, make_padding_mask +from .toucan import cache_toucan_examples, get_toucan_batches, load_toucan_contrastive_data +from .train_utils import ( + count_params, + create_config_from_args, + create_train_state, + load_checkpoint, + quantize_params, + save_checkpoint, + shard_batch, + wsd_schedule, +) + + +_GROUP_SIZE = 32 +_TOOL_CONTRASTIVE_WEIGHT = 1.0 +_AUDIO_TEXT_CONTRASTIVE_WEIGHT = 1.0 + + +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 _stage1_loss_fn(state, params, mel, transcript_text, tgt_in, tgt_out, causal_mask, rng, loss_mask, + contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask): + pad_id = 0 + src_mask = make_mel_padding_mask(mel) + tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) + spec_rng, drop_rng = jax.random.split(rng) + logits, slot_div = 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={"specaugment": spec_rng, "dropout": drop_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) + div_loss = 1e-4 * slot_div + at_loss = _audio_text_contrastive_loss(state, params, transcript_text, mel) + tc_loss = _tool_contrastive_loss( + state, params, contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask + ) + return ce_loss + z_loss + div_loss + _AUDIO_TEXT_CONTRASTIVE_WEIGHT * at_loss + _TOOL_CONTRASTIVE_WEIGHT * tc_loss + + +def _train_step(state, ema_params, mel, transcript_text, tgt_in, tgt_out, causal_mask, rng, loss_mask, + contrastive_q, contrastive_tools, contrastive_labels, contrastive_tool_mask): + ema_decay = 0.999 + loss, grads = jax.value_and_grad( + lambda p: _stage1_loss_fn( + state, + p, + mel, + transcript_text, + tgt_in, + tgt_out, + causal_mask, + rng, + loss_mask, + contrastive_q, + contrastive_tools, + contrastive_labels, + contrastive_tool_mask, + ) + )(state.params) + grads = jax.lax.pmean(grads, axis_name="batch") + 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_train_step(): + return jax.pmap(_train_step, axis_name="batch", donate_argnums=(0, 1)) + + +def _make_val_loss_fn(apply_fn): + @jax.jit + def 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 val_loss_batch + + +def _evaluate_val_ppl(val_loss_fn, params, prepared, batch_size, max_mel_len, n_mels, + gcs_prefix, max_dec_len, max_eval_samples=None): + val_causal = make_causal_mask(max_dec_len) + total_loss = 0.0 + total_toks = 0.0 + seen = 0 + for batch in get_transcription_batches( + prepared, + batch_size, + max_mel_len=max_mel_len, + n_mels=n_mels, + gcs_prefix=gcs_prefix, + shuffle=False, + require_local_mels=True, + ): + if max_eval_samples is not None and seen >= max_eval_samples: + break + mel, _, tgt_in, tgt_out, lm = batch + vl, vt = val_loss_fn(params, mel, tgt_in, tgt_out, val_causal, lm) + total_loss += float(vl) + total_toks += float(vt) + seen += len(mel) + return float(math.exp(min(total_loss / max(total_toks, 1.0), 20.0))) + + +def pretrain(args): + global _GROUP_SIZE, _TOOL_CONTRASTIVE_WEIGHT, _AUDIO_TEXT_CONTRASTIVE_WEIGHT + _GROUP_SIZE = getattr(args, "group_size", 32) + _TOOL_CONTRASTIVE_WEIGHT = getattr(args, "tool_contrastive_weight", 1.0) + _AUDIO_TEXT_CONTRASTIVE_WEIGHT = getattr(args, "audio_text_contrastive_weight", 1.0) + + num_devices = jax.local_device_count() + use_wandb = getattr(args, "wandb", False) + if use_wandb: + import wandb + if wandb.run is None: + wandb.init(project="needle-stage1", config=vars(args)) + + speech_prefix = getattr(args, "speech_gcs_prefix", EMILIA_SPEECH_GCS_PREFIX) + speech_val_ratio = getattr(args, "speech_val_ratio", 0.01) + speech_max_samples = getattr(args, "max_speech_samples", None) or getattr(args, "max_samples", None) + + print(f"\n[1/4] Loading tokenizer...") + tokenizer = get_tokenizer(max_samples=args.max_samples) + + print(f"\n[2/4] Loading Emilia metadata...") + train_prepared = load_prepared_transcription_data( + "train", + args.max_enc_len, + args.max_dec_len, + gcs_prefix=speech_prefix, + val_ratio=speech_val_ratio, + max_samples=speech_max_samples, + mmap=True, + ) + val_prepared = load_prepared_transcription_data( + "val", + args.max_enc_len, + args.max_dec_len, + gcs_prefix=speech_prefix, + val_ratio=speech_val_ratio, + max_samples=getattr(args, "max_eval_samples", None), + mmap=True, + ) + if train_prepared is None or val_prepared is None: + train_rows = load_emilia_speech_metadata( + "train", + gcs_prefix=speech_prefix, + max_samples=speech_max_samples, + val_ratio=speech_val_ratio, + seed=args.seed, + ) + val_rows = load_emilia_speech_metadata( + "val", + gcs_prefix=speech_prefix, + max_samples=getattr(args, "max_eval_samples", None), + val_ratio=speech_val_ratio, + seed=args.seed, + ) + train_prepared = prepare_transcription_pairs(train_rows, tokenizer, args.max_enc_len, args.max_dec_len) + val_prepared = prepare_transcription_pairs(val_rows, tokenizer, args.max_enc_len, args.max_dec_len) + print(f" {len(train_rows):,} train / {len(val_rows):,} val speech examples") + else: + print( + f" loaded cached Stage 1 speech data: " + f"{len(train_prepared['text_inputs']):,} train / {len(val_prepared['text_inputs']):,} val" + ) + + print(f"\n[3/4] Loading Toucan contrastive data...") + toucan_config = getattr(args, "toucan_config", None) or "Kimi-K2" + toucan_path = cache_toucan_examples( + config=toucan_config, + split="train", + max_samples=getattr(args, "toucan_max_samples", None), + tokenizer=tokenizer, + max_text_len=args.max_enc_len, + ) + toucan_batch = load_toucan_contrastive_data(toucan_path, max_text_len=args.max_enc_len) + + print(f"\n[4/4] Building model...") + resume_checkpoint = getattr(args, "checkpoint", None) + if resume_checkpoint: + ckpt_params, config, _ = load_checkpoint(resume_checkpoint) + print(f" loaded {resume_checkpoint}") + else: + config = create_config_from_args(args, n_mels=getattr(args, "n_mels", 80)) + ckpt_params = None + + effective_batch_size = args.batch_size * 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)) + + np.random.seed(args.seed) + rng = jax.random.PRNGKey(args.seed) + rng, init_rng = jax.random.split(rng) + + scaled_lr = args.lr * num_devices + 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) + + val_loss_fn = _make_val_loss_fn(state.apply_fn) + p_train_step = _make_p_train_step() + + ema_params = jax.tree.map(jnp.copy, state.params) + state = jax_utils.replicate(state) + ema_params = jax_utils.replicate(ema_params) + + param_count = count_params(jax_utils.unreplicate(state).params) + print(f"\n Parameters {param_count:,}") + print(f" d_model {config.d_model}") + print(f" Heads {config.num_heads} ({config.num_kv_heads} KV)") + print(f" Layers {config.num_encoder_layers} enc / {config.num_decoder_layers} dec") + print(f" Speech data {speech_prefix}") + print(f" Batch {args.batch_size} x {num_devices} = {effective_batch_size}") + print(f" Total steps {total_steps:,}\n") + + os.makedirs(args.checkpoint_dir, exist_ok=True) + causal_mask = jnp.broadcast_to( + make_causal_mask(args.max_dec_len), + (num_devices, 1, args.max_dec_len, args.max_dec_len), + ) + adam_schedule = wsd_schedule(scaled_lr, total_steps, warmup_steps) + muon_schedule = wsd_schedule(muon_lr, total_steps, warmup_steps) + + global_step = 0 + last_val_ppl = None + eval_every = getattr(args, "eval_every", 1000) + + for epoch in range(args.epochs): + losses = [] + speech_iter = PrefetchIterator( + lambda: get_transcription_batches( + train_prepared, + effective_batch_size, + max_mel_len=args.max_mel_len, + n_mels=args.n_mels, + gcs_prefix=speech_prefix, + require_local_mels=True, + ), + prefetch=2, + ) + contrastive_iter = PrefetchIterator( + lambda: get_toucan_batches(*toucan_batch, effective_batch_size), + prefetch=2, + ) + pbar = tqdm(range(batches_per_epoch), desc=f"Stage 1 Epoch {epoch + 1}/{args.epochs}") + + for _ in pbar: + mel, text_enc, tgt_in, tgt_out, lm = next(speech_iter) + tc_q, tc_tools, tc_labels, tc_tool_mask = next(contrastive_iter) + + 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) + 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) + + rng, step_rng = jax.random.split(rng) + step_rngs = jax.random.split(step_rng, num_devices) + t0 = time.perf_counter() + state, ema_params, loss, grad_norm = p_train_step( + state, + ema_params, + mel_b, + text_enc_b, + tgt_in_b, + tgt_out_b, + causal_mask, + step_rngs, + lm_b, + tc_q_b, + tc_tools_b, + tc_labels_b, + tc_tool_mask_b, + ) + dt = time.perf_counter() - t0 + loss_val = float(loss[0]) + grad_norm_val = float(grad_norm[0]) + losses.append(loss_val) + global_step += 1 + + if global_step % eval_every == 0 or global_step == total_steps: + eval_params = jax_utils.unreplicate(ema_params) + last_val_ppl = _evaluate_val_ppl( + val_loss_fn, + eval_params, + val_prepared, + args.batch_size, + args.max_mel_len, + args.n_mels, + speech_prefix, + args.max_dec_len, + max_eval_samples=getattr(args, "max_eval_samples", None), + ) + + pbar.set_postfix( + loss=f"{loss_val:.4f}", + ppl=f"{last_val_ppl:.2f}" if last_val_ppl is not None else "?", + ) + + if use_wandb: + wandb.log( + { + "train/loss": loss_val, + "train/grad_norm": grad_norm_val, + "train/adam_lr": float(adam_schedule(global_step)), + "train/muon_lr": float(muon_schedule(global_step)), + "train/tokens_per_sec": effective_batch_size * 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 {}), + } + ) + + speech_iter.close() + contrastive_iter.close() + + eval_params = jax_utils.unreplicate(ema_params) + val_ppl = _evaluate_val_ppl( + val_loss_fn, + eval_params, + val_prepared, + args.batch_size, + args.max_mel_len, + args.n_mels, + speech_prefix, + args.max_dec_len, + max_eval_samples=getattr(args, "max_eval_samples", None), + ) + 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}") + print(f" Val ppl {val_ppl:.2f}") + print(f" Checkpoint {ckpt_path}\n") + + if use_wandb: + wandb.finish() + print("Stage 1 pretraining complete.") diff --git a/src/tokenize_data.py b/src/tokenize_data.py index 04a6b5b..cfef6e3 100644 --- a/src/tokenize_data.py +++ b/src/tokenize_data.py @@ -1,12 +1,14 @@ -"""Standalone tokenization pipeline: train tokenizer and pre-tokenize all data. +"""Stage 2 tokenization pipeline for tool-call text data. -Trains the SentencePiece tokenizer and tokenizes train + val splits for both -text and voice data, caching everything on GCS. Running fresh overwrites all -existing tokenizer + caches. +Retrains the local SentencePiece tokenizer, optionally updates the shared GCS +tokenizer, and prepares the train + val text tool-call caches used by Stage 2. +The old unified-dataset mel precompute path is intentionally not part of this +pipeline anymore. Usage: needle tokenize # full run needle tokenize --max-samples 1000 # dev/test + needle tokenize --overwrite-gcs-tokenizer needle tokenize --cleanup # delete local cache after GCS upload """ @@ -16,62 +18,124 @@ from .data import ( CACHE_DIR, - GCS_CACHE_PATH, GCS_TOKENIZER_PATH, + EMILIA_SPEECH_GCS_PREFIX, TOKENIZER_DIR, _cache_key, _save_cache_metadata, get_tokenizer, + load_emilia_speech_metadata, load_tool_calls, - precompute_mels, + prefetch_emilia_mels, + prepare_transcription_pairs, prepare_tool_call_pairs, + save_prepared_transcription_data, train_tokenizer, upload_tokenizer_to_gcs, ) +from .toucan import cache_toucan_examples -def _clear_gcs_caches(): - """Remove existing GCS cache and tokenizer files.""" - for path in [GCS_CACHE_PATH + "/*", GCS_TOKENIZER_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 + "*" + print(f"Clearing {path} ...") + subprocess.run( + ["gcloud", "storage", "rm", "-r", path], + capture_output=True, text=True, + ) def _clear_local_caches(): - """Remove local .data_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): - print("=== Clearing existing caches ===") - _clear_gcs_caches() + print("=== Clearing existing local caches ===") _clear_local_caches() + if getattr(args, "overwrite_gcs_tokenizer", False): + print("\n=== Clearing shared GCS tokenizer ===") + _clear_gcs_tokenizer() - print("\n=== Training tokenizer ===") + print("\n=== Training local tokenizer ===") train_tokenizer(max_samples=args.max_samples, force=True) - upload_tokenizer_to_gcs() + if getattr(args, "overwrite_gcs_tokenizer", False): + upload_tokenizer_to_gcs() + else: + print("Skipping shared GCS tokenizer upload") tokenizer = get_tokenizer() + print("\n=== Caching Toucan tool definitions for Stage 1 ===") + 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 Emilia transcription data for Stage 1 ===") + speech_prefix = getattr(args, "speech_gcs_prefix", EMILIA_SPEECH_GCS_PREFIX) + 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, + gcs_prefix=speech_prefix, + max_samples=speech_max_samples, + 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, + prepared, + args.max_enc_len, + args.max_dec_len, + gcs_prefix=speech_prefix, + val_ratio=speech_val_ratio, + split_max_samples=speech_max_samples, + ) + 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=== Tokenizing text data + precomputing mels ===") + print("\n=== Mirroring Emilia mel files for Stage 1 ===") + use_rsync = speech_max_samples is None + mirrored = prefetch_emilia_mels(stage1_mel_uris, gcs_prefix=speech_prefix, use_rsync=use_rsync) + print(f"Mirrored {mirrored:,} Emilia mel files into local cache") + + print("\n=== Tokenizing Stage 2 text tool-call data ===") max_enc_len = getattr(args, "max_enc_len", 256) max_dec_len = getattr(args, "max_dec_len", 1024) - n_mels = getattr(args, "n_mels", 80) - max_mel_len = getattr(args, "max_mel_len", 1024) batch_size = getattr(args, "batch_size", 5000) for split in ("train", "val"): print(f"\n--- {split} split ---") - ds, global_indices = load_tool_calls( + ds = load_tool_calls( 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, @@ -79,13 +143,12 @@ def tokenize(args): ) text_cache_id = _cache_key("toolcall", len(ds), max_enc_len, max_dec_len) - mel_cache_id = precompute_mels( - global_indices[kept_indices], n_mels=n_mels, max_mel_len=max_mel_len, - cache_id_prefix=split, batch_size=batch_size, - ) - - _save_cache_metadata(split, text_cache_id, mel_cache_id, len(kept_indices), - max_enc_len, max_dec_len, n_mels, max_mel_len) + _save_cache_metadata(split, text_cache_id, None, len(kept_indices), + max_enc_len, max_dec_len, None, None, + split_max_samples=args.max_samples, + shuffle_before_split=getattr(args, "shuffle_before_split", False), + 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/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 f7bfdab..84dd2db 100644 --- a/src/train.py +++ b/src/train.py @@ -1,303 +1,64 @@ -import argparse import math import os -import pickle import time -from typing import NamedTuple import jax import jax.numpy as jnp import numpy as np import optax -from tqdm import tqdm from flax import jax_utils -from flax.training import train_state +from tqdm import tqdm from .data import ( - get_batches, get_tokenizer, get_speech_batches, - load_prepared_data, load_prepared_mels, - load_example_with_audio, - PrefetchIterator, count_batches, + PrefetchIterator, count_batches, get_batches, get_tokenizer, load_prepared_data, + prepare_audio_toolcall_val, ) -from .model import ( - EncoderDecoderTransformer, - TransformerConfig, - make_causal_mask, - make_padding_mask, - make_mel_padding_mask, +from .model import EncoderDecoderTransformer, make_causal_mask, make_mel_padding_mask, make_padding_mask +from .train_utils import ( + count_params, + create_config_from_args, + create_train_state, + load_checkpoint, + quantize_params, + save_checkpoint, + shard_batch, + wsd_schedule, + zero_non_decoder_grads, ) -def _newton_schulz(G, steps=5): - """Approximate polar decomposition via Newton-Schulz iteration.""" - a, b, c = 3.4445, -4.7750, 2.0315 - orig_dtype = G.dtype - G = G.astype(jnp.float32) - X = G / (jnp.linalg.norm(G) + 1e-7) - transposed = G.shape[0] > G.shape[1] - if transposed: - X = X.T - for _ in range(steps): - A = X @ X.T - B = b * A + c * A @ A - X = a * X + B @ X - if transposed: - X = X.T - return X.astype(orig_dtype) - - -class MuonState(NamedTuple): - mu: optax.Updates - - -def scale_by_muon(momentum=0.95, ns_steps=5): - """Muon gradient transform: orthogonalize 2D+ grads, then Nesterov momentum.""" - - def init_fn(params): - return MuonState(mu=jax.tree.map(jnp.zeros_like, params)) - - def update_fn(updates, state, params=None): - del params - - def ortho(g): - if g.ndim == 2: - return _newton_schulz(g, steps=ns_steps) - return g - - ortho_g = jax.tree.map(ortho, updates) - new_mu = jax.tree.map(lambda m, g: momentum * m + g, state.mu, ortho_g) - new_updates = jax.tree.map( - lambda g, m: g + momentum * m, ortho_g, new_mu - ) - return new_updates, MuonState(mu=new_mu) - - return optax.GradientTransformation(init_fn, update_fn) - - -def _param_labels(params): - """Label each param: 'muon' for Dense kernels, 'adam' for the rest.""" - - def _label(path, leaf): - name = path[-1].key if hasattr(path[-1], "key") else str(path[-1]) - if name == "kernel" and leaf.ndim == 2: - return "muon" - return "adam" - - return jax.tree_util.tree_map_with_path(_label, params) - - -def _wsd_schedule(peak_value, total_steps, warmup_steps, decay_ratio=0.15): - """Warmup-Stable-Decay schedule: linear warmup, hold peak, linear decay.""" - decay_steps = max(1, int(total_steps * decay_ratio)) - stable_steps = total_steps - warmup_steps - decay_steps - return optax.join_schedules( - [ - optax.linear_schedule(0.0, peak_value, warmup_steps), - optax.constant_schedule(peak_value), - optax.linear_schedule(peak_value, peak_value * 0.1, decay_steps), - ], - boundaries=[warmup_steps, warmup_steps + stable_steps], - ) - - -def create_train_state(rng, config, learning_rate, muon_lr, total_steps, warmup_steps): - model = EncoderDecoderTransformer(config) - - rng, init_rng = jax.random.split(rng) - dummy_src = jnp.ones((1, 128), dtype=jnp.int32) - dummy_tgt = jnp.ones((1, 128), dtype=jnp.int32) - dummy_mel = jnp.ones((1, 128, config.n_mels), dtype=jnp.float32) - variables = model.init( - {"params": init_rng}, - dummy_src, - dummy_tgt, - dummy_mel, - method="init_all", - ) - - adam_schedule = _wsd_schedule(learning_rate, total_steps, warmup_steps) - muon_schedule = _wsd_schedule(muon_lr, total_steps, warmup_steps) - - muon_opt = optax.chain( - scale_by_muon(momentum=0.95, ns_steps=5), - optax.add_decayed_weights(weight_decay=0.01), - optax.scale_by_schedule(muon_schedule), - optax.scale(-1.0), - ) - adam_opt = optax.chain( - optax.adamw(adam_schedule, b2=0.95, weight_decay=0.0), - ) - - tx = optax.chain( - optax.clip_by_global_norm(1.0), - optax.multi_transform( - {"muon": muon_opt, "adam": adam_opt}, - _param_labels, - ), - ) - return train_state.TrainState.create( - apply_fn=model.apply, - params=variables["params"], - tx=tx, - ) - - -def _fake_quantize_int4(w, group_size=32): - """Symmetric group-wise INT4 fake quantization with STE. - - Divides the input dimension (axis 0) into groups of `group_size` elements, - each with its own scale factor. Falls back to per-channel if in_features < group_size. - """ - in_feat, out_feat = w.shape - gs = min(group_size, in_feat) - - pad = (gs - in_feat % gs) % gs - if pad > 0: - w_padded = jnp.pad(w, ((0, pad), (0, 0))) - else: - w_padded = w - - num_groups = w_padded.shape[0] // gs - w_grouped = w_padded.reshape(num_groups, gs, out_feat) - - scale = jnp.max(jnp.abs(w_grouped), axis=1, keepdims=True) / 7.0 - scale = jnp.maximum(scale, 1e-8) - w_q = jnp.clip(jnp.round(w_grouped / scale), -8, 7) * scale - - w_q = w_q.reshape(-1, out_feat)[:in_feat] - - return w + jax.lax.stop_gradient(w_q - w) - - -def _cubic_sparsity_schedule(step, t_start, t_end, s_final): - """Cubic sparsity ramp (Zhu & Gupta 2017). Returns target sparsity at *step*.""" - if step < t_start: - return 0.0 - if step >= t_end: - return s_final - frac = (step - t_start) / (t_end - t_start) - return s_final * (1.0 - (1.0 - frac) ** 3) - - -def _make_prune_mask(params, sparsity, group_size): - """Compute block-prune mask entirely on-device (no numpy round-trips). - - Returns binary mask tree matching param shapes and dtypes. - """ - all_scores = [] - for _, leaf in jax.tree_util.tree_leaves_with_path(params): - if leaf.ndim != 2: - continue - in_feat, out_feat = leaf.shape - gs = min(group_size, in_feat) - pad = (gs - in_feat % gs) % gs - w = jnp.pad(leaf, ((0, pad), (0, 0))) if pad else leaf - scores = jnp.sum(jnp.abs(w.reshape(-1, gs, out_feat)), axis=1).ravel() - all_scores.append(scores) - - if not all_scores: - return jax.tree.map(jnp.ones_like, params) - - threshold = jnp.percentile(jnp.concatenate(all_scores), sparsity * 100) - - def _leaf_mask(path, leaf): - if leaf.ndim != 2: - return jnp.ones_like(leaf) - in_feat, out_feat = leaf.shape - gs = min(group_size, in_feat) - pad = (gs - in_feat % gs) % gs - w = jnp.pad(leaf, ((0, pad), (0, 0))) if pad else leaf - w_grouped = w.reshape(-1, gs, out_feat) - block_keep = (jnp.sum(jnp.abs(w_grouped), axis=1, keepdims=True) > threshold) - return jnp.broadcast_to(block_keep, w_grouped.shape).reshape(-1, out_feat)[:in_feat].astype(leaf.dtype) - - return jax.tree_util.tree_map_with_path(_leaf_mask, params) - - -def _quantize_params(params, group_size=32): - """Fake-quantize all Dense kernels in the param tree.""" - def _maybe_quantize(path, leaf): - name = path[-1].key if hasattr(path[-1], "key") else str(path[-1]) - if name == "kernel" and leaf.ndim == 2: - return _fake_quantize_int4(leaf, group_size=group_size) - return leaf - return jax.tree_util.tree_map_with_path(_maybe_quantize, params) - _GROUP_SIZE = 32 -_MAT_FACTORS = () -_MAT_FF_WIDTHS = () -_D_FF = 2048 -def _text_loss_fn(state, params, src, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask): +def _text_loss_fn(state, params, src, tgt_in, tgt_out, causal_mask, rng, loss_mask): pad_id = 0 src_mask = make_padding_mask(src, pad_id) tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) logits, slot_div = state.apply_fn( - {"params": _quantize_params(params, group_size=_GROUP_SIZE)}, - src, tgt_in, src_mask=src_mask, tgt_mask=tgt_mask, - ffn_mask=ffn_mask, + {"params": quantize_params(params, group_size=_GROUP_SIZE)}, + src, + tgt_in, + src_mask=src_mask, + tgt_mask=tgt_mask, deterministic=False, method="forward_masked", rngs={"dropout": rng}, ) 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 * loss_mask) / jnp.maximum(jnp.sum(loss_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 -def _speech_loss_fn(state, params, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask): - pad_id = 0 - src_mask = make_mel_padding_mask(mel) - tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) - spec_rng, drop_rng = jax.random.split(rng) - logits, slot_div = state.apply_fn( - {"params": _quantize_params(params, group_size=_GROUP_SIZE)}, - mel, tgt_in, src_mask=src_mask, tgt_mask=tgt_mask, - ffn_mask=ffn_mask, - deterministic=False, - method="forward_speech_masked", - rngs={"specaugment": spec_rng, "dropout": drop_rng}, - ) - 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) - 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 - - -def _make_ffn_mask(batch_size, d_ff, mat_ff_widths): - """Build (batch, d_ff) prefix mask: batch split equally across widths. - - First 1/N = full width (all ones), remaining 1/N sections = mat widths. - """ - n_widths = 1 + len(mat_ff_widths) # full + mat widths - per_width = batch_size // n_widths - arange = jnp.arange(d_ff) - rows = [jnp.ones((per_width, d_ff), dtype=jnp.bfloat16)] # full width - for k in mat_ff_widths: - rows.append((arange[None, :] < k).astype(jnp.bfloat16).repeat(per_width, axis=0)) - # Handle remainder (assign to full width) - remainder = batch_size - per_width * n_widths - if remainder > 0: - rows.append(jnp.ones((remainder, d_ff), dtype=jnp.bfloat16)) - 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(state, ema_params, src, tgt_in, tgt_out, causal_mask, rng, loss_mask): 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, 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) @@ -305,65 +66,48 @@ 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): - """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) - )(state.params) - grads = jax.lax.pmean(grads, axis_name="batch") - loss = jax.lax.pmean(loss, axis_name="batch") - grad_norm = optax.global_norm(grads) - state = state.apply_gradients(grads=grads) - masked_params = jax.tree.map(lambda w, m: w * m, state.params, prune_mask) - state = state.replace(params=masked_params) - ema_params = jax.tree.map(lambda e, p: ema_decay * e + (1 - ema_decay) * p, ema_params, masked_params) - return state, ema_params, loss, grad_norm +def _make_p_train_step(): + return jax.pmap(_train_step, axis_name="batch", donate_argnums=(0, 1)) -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) - )(state.params) - grads = jax.lax.pmean(grads, axis_name="batch") - 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 _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 _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.""" +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_loss_fn(state, p, mel, tgt_in, tgt_out, causal_mask, ffn_mask, rng, loss_mask) + 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) - masked_params = jax.tree.map(lambda w, m: w * m, state.params, prune_mask) - state = state.replace(params=masked_params) - ema_params = jax.tree.map(lambda e, p: ema_decay * e + (1 - ema_decay) * p, ema_params, masked_params) + 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_train_step(): - return jax.pmap(_train_step_text, axis_name="batch", donate_argnums=(0, 1)) - - -def _make_p_train_step_masked(): - return jax.pmap(_train_step_text_masked, axis_name="batch", donate_argnums=(0, 1)) - - -def _make_p_train_step_speech(): - return jax.pmap(_train_step_speech, axis_name="batch", donate_argnums=(0, 1)) - - -def _make_p_train_step_speech_masked(): - return jax.pmap(_train_step_speech_masked, axis_name="batch", donate_argnums=(0, 1)) +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): @@ -372,107 +116,110 @@ def val_loss_batch(params, src, tgt_in, tgt_out, causal_mask, loss_mask): pad_id = 0 src_mask = make_padding_mask(src, pad_id) tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) - logits = apply_fn( - {"params": params}, src, tgt_in, - src_mask=src_mask, tgt_mask=tgt_mask, - ) + logits = apply_fn({"params": params}, src, tgt_in, src_mask=src_mask, tgt_mask=tgt_mask) 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 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 - def val_loss_batch(params, src, tgt_in, tgt_out, causal_mask, loss_mask): - pad_id = 0 - src_mask = make_padding_mask(src, pad_id) - tgt_mask = causal_mask & make_padding_mask(tgt_in, pad_id) - logits, _, mat_logits = apply_fn( - {"params": params}, src, tgt_in, - src_mask=src_mask, tgt_mask=tgt_mask, - mat_ff_widths=(ff_width,), - method="forward_with_aux", - ) - trunc_logits = mat_logits[0].astype(jnp.float32) - loss = optax.softmax_cross_entropy_with_integer_labels(trunc_logits, tgt_out) - return jnp.sum(loss * loss_mask), jnp.sum(loss_mask) - return val_loss_batch +def _evaluate_val_ppl(val_loss_fn, params, val_enc, val_dec_in, val_dec_tgt, val_loss_mask, + batch_size, max_dec_len, max_eval_samples=None): + val_causal = make_causal_mask(max_dec_len) + total_loss = 0.0 + total_toks = 0.0 + seen = 0 + for batch in get_batches(val_enc, val_dec_in, val_dec_tgt, batch_size, shuffle=False, loss_mask=val_loss_mask): + if max_eval_samples is not None and seen >= max_eval_samples: + break + src, dec_in, dec_tgt, lm = batch + vl, vt = val_loss_fn(params, src, dec_in, dec_tgt, val_causal, lm) + total_loss += float(vl) + total_toks += float(vt) + seen += len(src) + 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 val_loss_batch(params, mel, tgt_in, tgt_out, causal_mask, loss_mask): + 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, + logits, _ = apply_fn( + {"params": params}, + mel, + tgt_in, + src_mask=src_mask, + tgt_mask=tgt_mask, deterministic=True, - method="forward_speech_with_aux", + method="forward_speech_masked", ) loss = optax.softmax_cross_entropy_with_integer_labels(logits.astype(jnp.float32), tgt_out) - 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. + return jnp.sum(loss * loss_mask), jnp.sum(loss_mask) - matryoshka_factor: how many times smaller the FFN widths are (e.g. 2 = half width). - Embedding and attention weights are unchanged. - """ - d = config.d_model - v = config.vocab_size - n_enc = config.num_encoder_layers - n_dec = config.num_decoder_layers - kv_dim = config.num_kv_heads * (d // config.num_heads) - m = config.num_memory_slots - d_ff = config.d_ff // matryoshka_factor + return speech_val_loss_batch - emb = v * d - attn = d * d + d * kv_dim * 2 + d * d - ffn = d * d_ff * 3 - mixer_token = m * d_ff * 3 - mixer_channel = d * d_ff * 3 - enc_block = attn + ffn + mixer_token + mixer_channel - dec_block = attn * 2 + ffn - total = emb + n_enc * enc_block + n_dec * dec_block - return int(total) +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. -def shard_batch(batch, num_devices): - """Reshape a batch array so leading dim is (num_devices, per_device_batch, ...).""" - return batch.reshape(num_devices, -1, *batch.shape[1:]) + 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): - num_devices = jax.local_device_count() - no_speech = getattr(args, "no_speech", False) - n_mels = getattr(args, "n_mels", 80) - max_mel_len = getattr(args, "max_mel_len", 1024) + global _GROUP_SIZE + _GROUP_SIZE = getattr(args, "group_size", 32) + num_devices = jax.local_device_count() use_wandb = getattr(args, "wandb", False) if use_wandb: import wandb if wandb.run is None: - wandb.init(project="needle-v1", config=vars(args)) - - total_data_steps = 4 if not no_speech else 3 - step_idx = 0 + wandb.init(project="needle-stage2", config=vars(args)) - step_idx += 1 - print(f"\n[{step_idx}/{total_data_steps}] Detecting devices...") - print(f" {num_devices} device(s) for data-parallel training") - - step_idx += 1 - print(f"\n[{step_idx}/{total_data_steps}] Loading tokenizer...") + print(f"\n[1/3] Loading tokenizer...") tokenizer = get_tokenizer(max_samples=args.max_samples) + del tokenizer - step_idx += 1 - print(f"\n[{step_idx}/{total_data_steps}] Loading prepared data from disk (mmap)...") + print(f"\n[2/3] Loading prepared tool-call data...") train_data = load_prepared_data("train", mmap=True) val_data = load_prepared_data("val", mmap=True) enc_inputs = train_data["enc_inputs"] @@ -483,593 +230,247 @@ def train(args): val_dec_in = val_data["dec_inputs"] val_dec_tgt = val_data["dec_targets"] val_loss_mask = val_data["loss_mask"] - print(f" {len(enc_inputs):,} train / {len(val_enc):,} val tool-call pairs (memory-mapped)") - - train_mels = None - val_mels = None - if not no_speech: - step_idx += 1 - print(f"\n[{step_idx}/{total_data_steps}] Loading precomputed mel spectrograms (mmap)...") - train_mels = load_prepared_mels(train_data["mel_cache_id"], mmap=True) - val_mels = load_prepared_mels(val_data["mel_cache_id"], mmap=True) - print(f" {len(train_mels):,} train / {len(val_mels):,} val mel spectrograms (memory-mapped)") - - effective_batch_size = args.batch_size * num_devices + 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 " + 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) if resume_checkpoint: - print(f"Resuming from checkpoint: {resume_checkpoint}") - with open(resume_checkpoint, "rb") as f: - ckpt_data = pickle.load(f) - ckpt_params = jax.tree.map(jnp.array, ckpt_data["params"]) - config = TransformerConfig(**ckpt_data["config"]) - print(f" Config: d={config.d_model}, heads={config.num_heads}, layers={config.num_encoder_layers}/{config.num_decoder_layers}") + ckpt_params, config, _ = load_checkpoint(resume_checkpoint) + print(f" loaded {resume_checkpoint}") else: - config = TransformerConfig( - d_model=args.d_model, - num_heads=args.num_heads, - num_kv_heads=getattr(args, "num_kv_heads", None) or args.num_heads, - num_encoder_layers=args.num_layers, - num_decoder_layers=getattr(args, "num_dec_layers", args.num_layers), - d_ff=getattr(args, "d_ff", None) or args.d_model * 4, - max_seq_len=max(args.max_enc_len, args.max_dec_len), - dtype=args.dtype, - activation=getattr(args, "activation", "drelu"), - num_memory_slots=getattr(args, "num_memory_slots", 64), - n_mels=n_mels, - dropout_rate=getattr(args, "dropout", 0.1), - ) + config = create_config_from_args(args, n_mels=getattr(args, "n_mels", 80)) + ckpt_params = None - global _GROUP_SIZE, _MAT_FACTORS, _MAT_FF_WIDTHS, _D_FF - _GROUP_SIZE = getattr(args, "group_size", 32) - _D_FF = config.d_ff - 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) - _MAT_FF_WIDTHS = tuple(config.d_ff // f for f in _MAT_FACTORS) - else: - _MAT_FACTORS = () - _MAT_FF_WIDTHS = () - n_widths = 1 + len(_MAT_FF_WIDTHS) if _MAT_FF_WIDTHS else 1 - p_train_step = _make_p_train_step() - p_train_step_masked = _make_p_train_step_masked() - p_train_step_speech = _make_p_train_step_speech() - p_train_step_speech_masked = _make_p_train_step_speech_masked() + effective_batch_size = args.batch_size * num_devices + batches_per_epoch = count_batches(len(enc_inputs), effective_batch_size) + total_steps = batches_per_epoch * args.epochs + warmup_steps = max(1, int(total_steps * args.warmup_ratio)) np.random.seed(args.seed) rng = jax.random.PRNGKey(args.seed) rng, init_rng = jax.random.split(rng) - mat_shared_input = getattr(args, "mat_shared_input", False) - unique_batch_size = effective_batch_size // n_widths if (mat_shared_input and n_widths > 1) else effective_batch_size - text_batches_per_epoch = count_batches(len(enc_inputs), unique_batch_size) - if not no_speech and train_mels is not None: - speech_batches_per_epoch = text_batches_per_epoch - else: - speech_batches_per_epoch = 0 - num_batches = text_batches_per_epoch + speech_batches_per_epoch - total_steps = num_batches * args.epochs - warmup_steps = max(1, int(total_steps * args.warmup_ratio)) - scaled_lr = args.lr * num_devices 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) - val_loss_fn = _make_val_loss_fn(state.apply_fn) - speech_vl_fn = _make_speech_val_loss_fn(state.apply_fn) if not no_speech else None + if ckpt_params is not None: + 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) - if resume_checkpoint: - state = state.replace(params=ckpt_params) - print(f" Loaded checkpoint params into train state") + 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) ema_params = jax_utils.replicate(ema_params) - param_count = sum(x.size for x in jax.tree.leaves(jax_utils.unreplicate(state).params)) - decay_steps = max(1, int(total_steps * 0.15)) - stable_steps = total_steps - warmup_steps - decay_steps - - print(f"\n ─────────────────────────────────────") - print(f" Parameters {param_count:>12,}") - print(f" d_model {config.d_model:>12}") - print(f" Heads {config.num_heads:>7} ({config.num_kv_heads} KV)") - print(f" Layers {config.num_encoder_layers:>7} enc / {config.num_decoder_layers} dec") - print(f" Memory slots {config.num_memory_slots:>12}") - print(f" Activation {config.activation:>12}") - print(f" Dtype {config.dtype:>12}") - print(f" Dropout {config.dropout_rate:>12}") - 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}") - print(f" max_mel_len {max_mel_len:>12}") - else: - print(f" Speech disabled") - print(f" ─────────────────────────────────────") - print(f" Devices {num_devices:>12}") - print(f" Batch {args.batch_size:>7} x {num_devices} = {effective_batch_size}") - print(f" Adam LR {args.lr:>7} x {num_devices} = {scaled_lr}") - print(f" Muon LR {args.muon_lr:>7.4f} -> {muon_lr:.4f}") - print(f" Schedule {warmup_steps}w / {stable_steps}s / {decay_steps}d (WSD)") - print(f" Total steps {total_steps:>12,}") - print(f" Epochs {args.epochs:>12}") - print(f" ─────────────────────────────────────\n") + param_count = count_params(jax_utils.unreplicate(state).params) + print(f"\n Parameters {param_count:,}") + print(f" d_model {config.d_model}") + print(f" Heads {config.num_heads} ({config.num_kv_heads} KV)") + print(f" Layers {config.num_encoder_layers} enc / {config.num_decoder_layers} dec") + print(f" Encoder frozen during Stage 2") + print(f" Batch {args.batch_size} x {num_devices} = {effective_batch_size}") + print(f" Total steps {total_steps:,}\n") os.makedirs(args.checkpoint_dir, exist_ok=True) - global_step = 0 causal_mask = jnp.broadcast_to( make_causal_mask(args.max_dec_len), (num_devices, 1, args.max_dec_len, args.max_dec_len), ) + adam_schedule = wsd_schedule(scaled_lr, total_steps, warmup_steps) + muon_schedule = wsd_schedule(muon_lr, total_steps, warmup_steps) - text_ffn_mask = _make_ffn_mask(args.batch_size, config.d_ff, _MAT_FF_WIDTHS) - text_ffn_mask = jnp.broadcast_to( - text_ffn_mask[None, :, :], - (num_devices, args.batch_size, config.d_ff), - ) - if n_widths > 1: - print(f" Mat factors {n_widths} (full + {', '.join(str(f)+'x' for f in _MAT_FACTORS)})") - if mat_shared_input: - print(f" Mat mode shared input ({unique_batch_size // num_devices}/dev x {n_widths} repeats)") - else: - print(f" Mat mode unique input ({args.batch_size}/dev, random width)") - - adam_schedule = _wsd_schedule(scaled_lr, total_steps, warmup_steps) - muon_schedule = _wsd_schedule(muon_lr, total_steps, warmup_steps) - tokens_per_batch = effective_batch_size * (args.max_enc_len + args.max_dec_len) - - eval_model = EncoderDecoderTransformer(config) - + global_step = 0 last_val_ppl = None - sparsity_ratio = getattr(args, "sparsity_ratio", 0.0) - prune_mask = None - gradual_sparsify_done = False - - 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) + last_nonempty_ppl = None + last_speech_ppl = None + best_nonempty_ppl = float("inf") + eval_every = getattr(args, "eval_every", 1000) - weight_prune_epoch = 0 if sparsity_ratio > 0 else -1 + # Speech replay (optional, disabled by default) + speech_replay_iter = None for epoch in range(args.epochs): - if epoch == weight_prune_epoch and not gradual_sparsify_done: - t_start = int(num_batches * prune_start_frac) - t_end = int(num_batches * prune_end_frac) - print(f"\nGradual magnitude sparsification: 0% -> {sparsity_ratio*100:.0f}% over epoch {epoch+1} " - f"(steps {t_start}-{t_end}/{num_batches}, interval={prune_interval}, group_size={_GROUP_SIZE})") - epoch_step = 0 - - text_losses = [] - speech_losses = [] - text_batch_iter = PrefetchIterator( - lambda: get_batches(enc_inputs, dec_inputs, dec_targets, unique_batch_size, - loss_mask=train_loss_mask), + losses = [] + batch_iter = PrefetchIterator( + lambda: get_batches(enc_inputs, dec_inputs, dec_targets, effective_batch_size, loss_mask=train_loss_mask), prefetch=4, ) + pbar = tqdm(range(batches_per_epoch), desc=f"Stage 2 Epoch {epoch + 1}/{args.epochs}") - speech_batch_iter = None - if not no_speech and train_mels is not None: - speech_batch_iter = PrefetchIterator( - lambda: get_speech_batches(train_mels, dec_inputs, dec_targets, unique_batch_size, - loss_mask=train_loss_mask), - prefetch=4, - ) - - steps_this_epoch = text_batches_per_epoch + speech_batches_per_epoch - text_idx = 0 - speech_idx = 0 - speech_loss_val = None - pbar = tqdm(range(steps_this_epoch), desc=f"Epoch {epoch + 1}/{args.epochs}") + for _ in pbar: + src, tgt_in, tgt_out, lm = next(batch_iter) + 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) - for step_i in pbar: + rng, step_rng = jax.random.split(rng) + step_rngs = jax.random.split(step_rng, num_devices) t0 = time.perf_counter() - - do_speech = (step_i % 2 == 1) and speech_idx < speech_batches_per_epoch - do_text = not do_speech and text_idx < text_batches_per_epoch - if not do_speech and not do_text: - if text_idx < text_batches_per_epoch: - do_text = True - elif speech_idx < speech_batches_per_epoch: - do_speech = True - else: - break - - step_grad_norm = None - - if do_text: - src, tgt_in, tgt_out, lm = next(text_batch_iter) - text_idx += 1 - - if n_widths > 1 and mat_shared_input: - per_width = args.batch_size // n_widths - def _tile_for_mat(arr): - s = arr.reshape(num_devices, per_width, *arr.shape[1:]) - return np.tile(s, (1, n_widths) + (1,) * (arr.ndim - 1)) - src_b = _tile_for_mat(src) - tgt_in_b = _tile_for_mat(tgt_in) - tgt_out_b = _tile_for_mat(tgt_out) - lm_b = _tile_for_mat(lm) - 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) - - 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, - ) - 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, - ) - - text_loss_val = float(loss[0]) - text_losses.append(text_loss_val) - step_grad_norm = float(grad_norm[0]) - global_step += 1 - - else: - mel_batch, sp_tgt_in, sp_tgt_out, sp_lm = next(speech_batch_iter) - speech_idx += 1 - - if n_widths > 1 and mat_shared_input: - per_width = args.batch_size // n_widths - def _tile_sp(arr): - s = arr.reshape(num_devices, per_width, *arr.shape[1:]) - return np.tile(s, (1, n_widths) + (1,) * (arr.ndim - 1)) - mel_b = _tile_sp(mel_batch) - sp_tgt_in_b = _tile_sp(sp_tgt_in) - sp_tgt_out_b = _tile_sp(sp_tgt_out) - sp_lm_b = _tile_sp(sp_lm) - else: - mel_b = shard_batch(mel_batch, num_devices) + state, ema_params, loss, grad_norm = p_train_step( + state, ema_params, src_b, tgt_in_b, tgt_out_b, causal_mask, step_rngs, lm_b + ) + dt = time.perf_counter() - t0 + loss_val = float(loss[0]) + grad_norm_val = float(grad_norm[0]) + 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, spec_rng = jax.random.split(rng) - spec_rngs = jax.random.split(spec_rng, num_devices) - - 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, + 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 ) - 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, - ) - speech_loss_val = float(sp_loss[0]) - speech_losses.append(speech_loss_val) - step_grad_norm = float(sp_grad_norm[0]) - text_loss_val = text_losses[-1] if text_losses else float("nan") - global_step += 1 - - if epoch == weight_prune_epoch and not gradual_sparsify_done: - epoch_step += 1 - current_sparsity = _cubic_sparsity_schedule(epoch_step, t_start, t_end, sparsity_ratio) - if epoch_step >= t_start and epoch_step % prune_interval == 0 and current_sparsity > 0: - ema_unr = jax_utils.unreplicate(ema_params) - mask = _make_prune_mask(ema_unr, current_sparsity, _GROUP_SIZE) - del ema_unr - prune_mask = jax_utils.replicate(mask) - del mask + except StopIteration: + pass - dt = time.perf_counter() - t0 - eval_every = getattr(args, "eval_every", 100) if global_step % eval_every == 0 or global_step == total_steps: - _eval_params = jax_utils.unreplicate(ema_params) - val_causal = make_causal_mask(args.max_dec_len) - total_loss, total_toks = 0.0, 0.0 - for vb in get_batches(val_enc, val_dec_in, val_dec_tgt, args.batch_size, shuffle=False, loss_mask=val_loss_mask): - vl, vt = val_loss_fn(_eval_params, vb[0], vb[1], vb[2], val_causal, vb[3]) - total_loss += float(vl) - total_toks += float(vt) - last_val_ppl = float(math.exp(min(total_loss / max(total_toks, 1), 20))) - - del _eval_params - - postfix = { - "speech_loss": f"{speech_loss_val:.4f}" if speech_loss_val is not None else "-", - "text_loss": f"{text_loss_val:.4f}", - "text_ppl": f"{last_val_ppl:.2f}" if last_val_ppl is not None else "?", - } - if sparsity_ratio > 0: - if epoch == weight_prune_epoch and not gradual_sparsify_done: - postfix["sparsification"] = f"{current_sparsity*100:.1f}%" - else: - postfix["sparsification"] = "done" - pbar.set_postfix(**postfix) + eval_params = jax_utils.unreplicate(ema_params) + last_val_ppl = _evaluate_val_ppl( + val_loss_fn, + eval_params, + val_enc, + val_dec_in, + val_dec_tgt, + val_loss_mask, + args.batch_size, + 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: - log_dict = { - "train/text_loss": text_loss_val, - "train/grad_norm": step_grad_norm, - "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, - } - 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: - log_dict["train/scheduled_sparsity"] = current_sparsity - if global_step % eval_every == 0 or global_step == total_steps: - log_dict["val/text_ppl"] = last_val_ppl - wandb.log(log_dict) - - text_batch_iter.close() - if speech_batch_iter is not None: - speech_batch_iter.close() - - if epoch == weight_prune_epoch and not gradual_sparsify_done: - gradual_sparsify_done = True - if prune_mask is None: - ema_unr = jax_utils.unreplicate(ema_params) - mask = _make_prune_mask(ema_unr, sparsity_ratio, _GROUP_SIZE) - del ema_unr - prune_mask = jax_utils.replicate(mask) - del mask - state = state.replace( - params=jax.tree.map(lambda w, m: w * m, state.params, prune_mask)) - ema_params = jax.tree.map(lambda w, m: w * m, ema_params, prune_mask) - final_pruned = jax.tree.map(np.array, jax_utils.unreplicate(ema_params)) - total_p = sum(x.size for x in jax.tree.leaves(final_pruned)) - zero_p = sum(int(np.sum(np.abs(x) < 1e-6)) for x in jax.tree.leaves(final_pruned)) - print(f"\n Gradual sparsification complete — mask locked.") - print(f" Final sparsity: {zero_p/total_p*100:.2f}% ({zero_p:,}/{total_p:,} near-zero)") - del final_pruned - - epoch_avg_loss = sum(text_losses) / len(text_losses) if text_losses else float("nan") - 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") + wandb.log( + { + "train/loss": loss_val, + "train/grad_norm": grad_norm_val, + "train/adam_lr": float(adam_schedule(global_step)), + "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, "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 {}), + } + ) + + batch_iter.close() eval_params = jax_utils.unreplicate(ema_params) - val_causal = make_causal_mask(args.max_dec_len) - - q_params = _quantize_params(eval_params, group_size=_GROUP_SIZE) - mat_vl_fns = {} - if _MAT_FACTORS: - _apply_fn = jax_utils.unreplicate(state).apply_fn - mat_vl_fns = {f: _make_mat_val_loss_fn(_apply_fn, fw) - for f, fw in zip(_MAT_FACTORS, _MAT_FF_WIDTHS)} - del _apply_fn - - full_loss, full_toks = 0.0, 0.0 - 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): - 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) - vl, vt = val_loss_fn(q_params, src, dec_in, dec_tgt, val_causal, lm) - q_loss += float(vl); q_toks += float(vt) - for f, fn in mat_vl_fns.items(): - vl, vt = fn(eval_params, src, dec_in, dec_tgt, val_causal, lm) - mat_accum[f][0] += float(vl) - mat_accum[f][1] += float(vt) - - 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 - - mat_results = {} - for f in _MAT_FACTORS: - avg = mat_accum[f][0] / max(mat_accum[f][1], 1) - mat_results[f] = (float(math.exp(min(avg, 20))), - _estimate_mat_params(config, f), config.d_ff // f) - - 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): - 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))) - - params_np = jax.tree.map(np.array, eval_params) - total_params = sum(x.size for x in jax.tree.leaves(params_np)) - 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" + val_ppl = _evaluate_val_ppl( + val_loss_fn, + eval_params, + val_enc, + val_dec_in, + val_dec_tgt, + val_loss_mask, + args.batch_size, + 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) - with open(ckpt_path, "wb") as f: - pickle.dump({"params": params_np, "config": config.__dict__}, f) - del params_np - - from .test import measure_throughput - from .run import generate, generate_from_audio - tp = measure_throughput(eval_model, eval_params, tokenizer, num_runs=5) + save_checkpoint(ckpt_path, eval_params, config, extra={"stage": "stage2"}) - _val_start = len(enc_inputs) - val_kept = val_data["kept_indices"] - - # Pick 5 display samples (shuffled for diversity): 4 with tool calls, 1 without - sample_rng = np.random.RandomState(epoch + 7) - sample_pool = sample_rng.permutation(len(val_kept)) - display_with, display_without = [], [] - for k in sample_pool: - if len(display_with) >= 4 and len(display_without) >= 1: - break - ds_idx = int(val_kept[k]) + _val_start - pair = load_example_with_audio(ds_idx) - is_empty = pair["answers"].strip() in ("", "[]") - if not is_empty and len(display_with) < 4: - display_with.append((ds_idx, pair)) - elif is_empty and len(display_without) < 1: - display_without.append((ds_idx, pair)) - display_pairs = display_with + display_without - - unified_samples = [] - for i, (ds_idx, pair) in enumerate(display_pairs): - text_pred = generate( - eval_model, eval_params, tokenizer, pair["query"], - tools=pair["tools"], max_gen_len=args.max_dec_len, seed=i, stream=False, - ).strip() - - 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, - ).strip() - - unified_samples.append({ - "query": pair["query"], - "tools": pair["tools"], - "ref": pair["answers"], - "text": text_pred, - "voice": voice_pred, - }) - - import json as _json_mod - tc_n, tc_exact, tc_name_tp, tc_name_fp, tc_name_fn = 0, 0, 0, 0, 0 - tc_call_tp, tc_call_fp, tc_call_fn, tc_parse_err = 0, 0, 0, 0 - tc_with_n = 25 - tc_without_n = 5 - tc_rng = np.random.RandomState(epoch + 42) - tc_pool = tc_rng.permutation(len(val_kept)) - tc_with, tc_without = [], [] - for k in tc_pool: - if len(tc_with) >= tc_with_n and len(tc_without) >= tc_without_n: - break - ds_idx = int(val_kept[k]) + _val_start - pair = load_example_with_audio(ds_idx) - ref_text = pair["answers"].strip() - is_empty = ref_text in ("", "[]") - if not is_empty and len(tc_with) < tc_with_n: - tc_with.append((ds_idx, pair)) - elif is_empty and len(tc_without) < tc_without_n: - tc_without.append((ds_idx, pair)) - tc_eval_pairs = tc_with + tc_without - - def _call_key(c): - if not isinstance(c, dict): return None - return _json_mod.dumps({"name": c.get("name"), "arguments": c.get("arguments")}, sort_keys=True) - - for i, (ds_idx, pair) in enumerate(tc_eval_pairs): - ref_text = pair["answers"].strip() - pred_text = generate( - eval_model, eval_params, tokenizer, pair["query"], - tools=pair["tools"], max_gen_len=args.max_dec_len, seed=ds_idx, stream=False, - ).strip() - try: - ref_calls = _json_mod.loads(ref_text) - except (ValueError, TypeError): - ref_calls = [] - try: - pred_calls = _json_mod.loads(pred_text) - if not isinstance(pred_calls, list): - pred_calls = [pred_calls] if isinstance(pred_calls, dict) else [] - except (ValueError, TypeError): - tc_parse_err += 1 - pred_calls = [] - tc_n += 1 - if _json_mod.dumps(pred_calls, sort_keys=True) == _json_mod.dumps(ref_calls, sort_keys=True): - tc_exact += 1 - ref_names = {c["name"] for c in ref_calls if isinstance(c, dict) and "name" in c} - pred_names = {c["name"] for c in pred_calls if isinstance(c, dict) and "name" in c} - tc_name_tp += len(pred_names & ref_names) - tc_name_fp += len(pred_names - ref_names) - tc_name_fn += len(ref_names - pred_names) - rk = {_call_key(c) for c in ref_calls} - {None} - pk = {_call_key(c) for c in pred_calls} - {None} - tc_call_tp += len(pk & rk) - tc_call_fp += len(pk - rk) - tc_call_fn += len(rk - pk) - - tc_metrics = {} - if tc_n > 0: - tc_metrics["parse_rate"] = 1.0 - tc_parse_err / tc_n - tc_metrics["exact_match"] = tc_exact / tc_n - np_ = tc_name_tp + tc_name_fp - nr_ = tc_name_tp + tc_name_fn - tc_metrics["name_f1"] = 2 * tc_name_tp / max(np_ + nr_, 1) - cp_ = tc_call_tp + tc_call_fp - cr_ = tc_call_tp + tc_call_fn - tc_metrics["call_f1"] = 2 * tc_call_tp / max(cp_ + cr_, 1) - - del eval_params - - final_speech_loss = speech_losses[-1] if speech_losses else None - print(f"\n ─────────────────────────────────────") - print(f" Epoch {epoch + 1}/{args.epochs}") - print(f" ─────────────────────────────────────") - print(f" Text loss {final_loss:>12.4f}") - print(f" Text val ppl {last_val_ppl:>12.2f}") - if final_speech_loss is not None: - print(f" Speech loss {final_speech_loss:>12.4f}") - if speech_val_ppl is not None: - print(f" Speech val ppl {speech_val_ppl:>12.2f}") - print(f" Quant val ppl {quant_val_ppl:>12.2f} (INT4 g{_GROUP_SIZE})") - print(f" Sparsity {sparsity:>11.2f}% ({near_zero:,}/{total_params:,})") - if mat_results: - print(f" ─────────────────────────────────────") - print(f" Matryoshka sub-models:") - print(f" {'factor':>6} {'d_ff':>6} {'val ppl':>10} {'params':>12}") - print(f" {'1x':>6} {config.d_ff:>6} {last_val_ppl:>10.2f} {total_params:>12,} (full)") - for factor in sorted(mat_results.keys()): - mat_ppl, mat_params, ff_w = mat_results[factor] - print(f" {str(factor)+'x':>6} {ff_w:>6} {mat_ppl:>10.2f} {mat_params:>12,}") - if tc_metrics: - print(f" ─── Tool-Call Accuracy ({tc_n} samples) ──") - print(f" JSON parse {tc_metrics['parse_rate']:>10.1%}") - print(f" Exact match {tc_metrics['exact_match']:>10.1%}") - print(f" Name F1 {tc_metrics['name_f1']:>10.3f}") - print(f" Call F1 {tc_metrics['call_f1']:>10.3f}") - print(f" ─────────────────────────────────────") - print(f" Throughput {tp['tokens_per_second']:>10.1f} tok/s") - print(f" Latency {tp['avg_latency_s']:>11.3f}s") - if unified_samples: - print(f" ─── Samples ({len(unified_samples)}) ───────────────────") - for j, s in enumerate(unified_samples): - print(f" [{j+1}] Query: {s['query'][:120]}") - tools_short = s["tools"][:120] - if len(s["tools"]) > 120: - tools_short += "..." - print(f" Tools: {tools_short}") - print(f" Ref: {s['ref'][:200] or '[]'}") - print(f" Text: {s['text'][:200] or '(empty)'}") - if s["voice"] is not None: - print(f" Voice: {s['voice'][:200] or '(empty)'}") - if j < len(unified_samples) - 1: - print() - print(f" ─────────────────────────────────────") - print(f" Checkpoint: {ckpt_path}") - print(f" ─────────────────────────────────────\n") - - if use_wandb: - log_dict = { - "epoch/text_loss": final_loss, - "epoch/text_val_ppl": last_val_ppl, - "epoch/quant_val_ppl": quant_val_ppl, - "epoch/weight_sparsity": sparsity, - "epoch": epoch + 1, - } - if final_speech_loss is not None: - log_dict["epoch/speech_loss"] = final_speech_loss - if speech_val_ppl is not None: - log_dict["epoch/speech_val_ppl"] = speech_val_ppl - for factor, (mat_ppl, mat_params, _) in mat_results.items(): - log_dict[f"epoch/mat_ppl_{factor}x"] = mat_ppl - log_dict[f"epoch/mat_params_{factor}x"] = mat_params - if tc_metrics: - log_dict["epoch/tc_parse_rate"] = tc_metrics["parse_rate"] - log_dict["epoch/tc_exact_match"] = tc_metrics["exact_match"] - log_dict["epoch/tc_name_f1"] = tc_metrics["name_f1"] - log_dict["epoch/tc_call_f1"] = tc_metrics["call_f1"] - wandb.log(log_dict) + 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} (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: wandb.finish() - print("\nTraining complete.") - - + print("Stage 2 training complete.") diff --git a/src/train_utils.py b/src/train_utils.py new file mode 100644 index 0000000..29dcfb1 --- /dev/null +++ b/src/train_utils.py @@ -0,0 +1,189 @@ +import math +import pickle +from typing import NamedTuple + +import jax +import jax.numpy as jnp +import numpy as np +import optax +from flax.training import train_state + +from .model import EncoderDecoderTransformer, TransformerConfig + + +def _newton_schulz(G, steps=5): + a, b, c = 3.4445, -4.7750, 2.0315 + orig_dtype = G.dtype + G = G.astype(jnp.float32) + X = G / (jnp.linalg.norm(G) + 1e-7) + transposed = G.shape[0] > G.shape[1] + if transposed: + X = X.T + for _ in range(steps): + A = X @ X.T + B = b * A + c * A @ A + X = a * X + B @ X + if transposed: + X = X.T + return X.astype(orig_dtype) + + +class MuonState(NamedTuple): + mu: optax.Updates + + +def scale_by_muon(momentum=0.95, ns_steps=5): + def init_fn(params): + return MuonState(mu=jax.tree.map(jnp.zeros_like, params)) + + def update_fn(updates, state, params=None): + del params + + def ortho(g): + if g.ndim == 2: + return _newton_schulz(g, steps=ns_steps) + return g + + ortho_g = jax.tree.map(ortho, updates) + new_mu = jax.tree.map(lambda m, g: momentum * m + g, state.mu, ortho_g) + new_updates = jax.tree.map(lambda g, m: g + momentum * m, ortho_g, new_mu) + return new_updates, MuonState(mu=new_mu) + + return optax.GradientTransformation(init_fn, update_fn) + + +def _param_labels(params): + def _label(path, leaf): + name = path[-1].key if hasattr(path[-1], "key") else str(path[-1]) + if name == "kernel" and leaf.ndim == 2: + return "muon" + return "adam" + + return jax.tree_util.tree_map_with_path(_label, params) + + +def wsd_schedule(peak_value, total_steps, warmup_steps, decay_ratio=0.15): + decay_steps = max(1, int(total_steps * decay_ratio)) + stable_steps = max(0, total_steps - warmup_steps - decay_steps) + return optax.join_schedules( + [ + optax.linear_schedule(0.0, peak_value, warmup_steps), + optax.constant_schedule(peak_value), + optax.linear_schedule(peak_value, peak_value * 0.1, decay_steps), + ], + boundaries=[warmup_steps, warmup_steps + stable_steps], + ) + + +def create_config_from_args(args, n_mels=80): + return TransformerConfig( + d_model=args.d_model, + num_heads=args.num_heads, + num_kv_heads=getattr(args, "num_kv_heads", None) or args.num_heads, + num_encoder_layers=args.num_layers, + num_decoder_layers=getattr(args, "num_dec_layers", args.num_layers), + d_ff=getattr(args, "d_ff", None) or args.d_model * 4, + max_seq_len=max(args.max_enc_len, args.max_dec_len), + dtype=args.dtype, + activation=getattr(args, "activation", "drelu"), + num_memory_slots=getattr(args, "num_memory_slots", 64), + n_mels=n_mels, + dropout_rate=getattr(args, "dropout", 0.0), + ) + + +def create_train_state(rng, config, learning_rate, muon_lr, total_steps, warmup_steps): + model = EncoderDecoderTransformer(config) + + rng, init_rng = jax.random.split(rng) + dummy_src = jnp.ones((1, 128), dtype=jnp.int32) + dummy_tgt = jnp.ones((1, 128), dtype=jnp.int32) + dummy_mel = jnp.ones((1, 128, config.n_mels), dtype=jnp.float32) + variables = model.init( + {"params": init_rng}, + dummy_src, + dummy_tgt, + dummy_mel, + method="init_all", + ) + + adam_schedule = wsd_schedule(learning_rate, total_steps, warmup_steps) + muon_schedule = wsd_schedule(muon_lr, total_steps, warmup_steps) + + muon_opt = optax.chain( + scale_by_muon(momentum=0.95, ns_steps=5), + optax.add_decayed_weights(weight_decay=0.01), + optax.scale_by_schedule(muon_schedule), + optax.scale(-1.0), + ) + adam_opt = optax.chain(optax.adamw(adam_schedule, b2=0.95, weight_decay=0.0)) + + tx = optax.chain( + optax.clip_by_global_norm(1.0), + optax.multi_transform( + {"muon": muon_opt, "adam": adam_opt}, + _param_labels, + ), + ) + return train_state.TrainState.create( + apply_fn=model.apply, + params=variables["params"], + tx=tx, + ) + + +def fake_quantize_int4(w, group_size=32): + in_feat, out_feat = w.shape + gs = min(group_size, in_feat) + pad = (gs - in_feat % gs) % gs + w_padded = jnp.pad(w, ((0, pad), (0, 0))) if pad else w + grouped = w_padded.reshape(-1, gs, out_feat) + scale = jnp.max(jnp.abs(grouped), axis=1, keepdims=True) / 7.0 + scale = jnp.maximum(scale, 1e-8) + w_q = jnp.clip(jnp.round(grouped / scale), -8, 7) * scale + w_q = w_q.reshape(-1, out_feat)[:in_feat] + return w + jax.lax.stop_gradient(w_q - w) + + +def quantize_params(params, group_size=32): + def _maybe_quantize(path, leaf): + name = path[-1].key if hasattr(path[-1], "key") else str(path[-1]) + if name == "kernel" and leaf.ndim == 2: + return fake_quantize_int4(leaf, group_size=group_size) + return leaf + + return jax.tree_util.tree_map_with_path(_maybe_quantize, params) + + +def shard_batch(batch, num_devices): + return batch.reshape(num_devices, -1, *batch.shape[1:]) + + +def count_params(params): + return sum(x.size for x in jax.tree.leaves(params)) + + +def save_checkpoint(path, params, config, extra=None): + payload = {"params": jax.tree.map(np.array, params), "config": config.__dict__} + if extra: + payload.update(extra) + with open(path, "wb") as f: + pickle.dump(payload, f) + + +def load_checkpoint(path): + with open(path, "rb") as f: + data = pickle.load(f) + params = jax.tree.map(jnp.array, data["params"]) + config = TransformerConfig(**data["config"]) + return params, config, data + + +def zero_non_decoder_grads(grads): + def _mask(path, leaf): + top = path[0].key if hasattr(path[0], "key") else str(path[0]) + if top == "decoder": + return leaf + return jnp.zeros_like(leaf) + + return jax.tree_util.tree_map_with_path(_mask, grads)