Spaces:
Sleeping
Sleeping
| """Train compound MIDI-text contrastive model from a CompoundGPT checkpoint.""" | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import math | |
| import random | |
| import sys | |
| import time | |
| from pathlib import Path | |
| from typing import Any, Dict, List, Optional, Sequence, Tuple | |
| import pretty_midi | |
| import torch | |
| from torch.optim import AdamW | |
| from torch.utils.data import DataLoader, Dataset | |
| _SCRIPT_DIR = Path(__file__).resolve().parent | |
| _ROOT = _SCRIPT_DIR.parent | |
| if str(_SCRIPT_DIR) not in sys.path: | |
| sys.path.insert(0, str(_SCRIPT_DIR)) | |
| from compound import SENTINELS, STEP_PAD, encode_compound # noqa: E402 | |
| from compound_model import ( # noqa: E402 | |
| CompoundGPT, | |
| CompoundGPTConfig, | |
| default_compound_config, | |
| ) | |
| from contrastive_model import CompoundMidiTextContrastiveModel # noqa: E402 | |
| def _pick_device() -> torch.device: | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| mps = getattr(torch.backends, "mps", None) | |
| if mps is not None and mps.is_available(): | |
| return torch.device("mps") | |
| return torch.device("cpu") | |
| def _load_jsonl_records(path: Path) -> List[Dict[str, Any]]: | |
| records: List[Dict[str, Any]] = [] | |
| with open(path) as f: | |
| for line in f: | |
| line = line.strip() | |
| if line: | |
| records.append(json.loads(line)) | |
| return records | |
| def _caption_from_record(record: Dict[str, Any]) -> str: | |
| llm_caption = str(record.get("caption", "")).strip() | |
| if llm_caption: | |
| return llm_caption | |
| return str(record.get("caption_template", "")).strip() | |
| def _path_from_record(record: Dict[str, Any]) -> str: | |
| for k in ("path", "midi_path", "midi"): | |
| value = str(record.get(k, "")).strip() | |
| if value: | |
| return value | |
| return "" | |
| class CompoundCaptionDataset(Dataset): | |
| """Dataset yielding padded compound steps + captions.""" | |
| def __init__( | |
| self, | |
| records: Sequence[Dict[str, Any]], | |
| is_train: bool, | |
| max_seq_len: int, | |
| seed: int, | |
| ) -> None: | |
| self.records = list(records) | |
| self.is_train = is_train | |
| self.max_seq_len = max_seq_len | |
| self._rng = random.Random(seed) | |
| self._cache: Dict[int, Optional[List[List[int]]]] = {} | |
| self._valid_indices: set[int] = set() | |
| def __len__(self) -> int: | |
| return len(self.records) | |
| def _encode_idx(self, idx: int) -> Optional[List[List[int]]]: | |
| if idx in self._cache: | |
| return self._cache[idx] | |
| rec = self.records[idx] | |
| path = _path_from_record(rec) | |
| if not path: | |
| self._cache[idx] = None | |
| return None | |
| try: | |
| pm = pretty_midi.PrettyMIDI(path) | |
| steps = encode_compound(pm) | |
| except Exception: | |
| self._cache[idx] = None | |
| return None | |
| if not steps: | |
| self._cache[idx] = None | |
| return None | |
| self._cache[idx] = steps | |
| self._valid_indices.add(idx) | |
| return steps | |
| def _crop(self, steps: List[List[int]]) -> List[List[int]]: | |
| if self.is_train: | |
| max_start = len(steps) - self.max_seq_len | |
| start = 0 if max_start <= 0 else self._rng.randint(0, max_start) | |
| return steps[start : start + self.max_seq_len] | |
| return steps[: self.max_seq_len] | |
| def _pad(self, steps: List[List[int]]) -> List[List[int]]: | |
| if len(steps) >= self.max_seq_len: | |
| return steps | |
| pad_step = list(SENTINELS) | |
| pad_step[0] = STEP_PAD | |
| return steps + [pad_step] * (self.max_seq_len - len(steps)) | |
| def __getitem__(self, idx: int) -> Dict[str, Any]: | |
| steps = self._encode_idx(idx) | |
| rec = self.records[idx] | |
| if steps is None: | |
| if not self._valid_indices: | |
| for j in range(len(self.records)): | |
| self._encode_idx(j) | |
| if not self._valid_indices: | |
| raise RuntimeError("No valid MIDI records for compound contrastive.") | |
| repl_idx = self._rng.choice(list(self._valid_indices)) | |
| steps = self._encode_idx(repl_idx) | |
| rec = self.records[repl_idx] | |
| assert steps is not None | |
| steps = self._pad(self._crop(steps)) | |
| return { | |
| "compound_input": torch.tensor(steps, dtype=torch.long), | |
| "caption": _caption_from_record(rec), | |
| "path": _path_from_record(rec), | |
| } | |
| def _collate(items: List[Dict[str, Any]]) -> Dict[str, Any]: | |
| return { | |
| "compound_input": torch.stack([x["compound_input"] for x in items], dim=0), | |
| "captions": [x["caption"] for x in items], | |
| "paths": [x["path"] for x in items], | |
| } | |
| def _split_records( | |
| records: List[Dict[str, Any]], split_ratio: float, seed: int | |
| ) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]: | |
| rng = random.Random(seed) | |
| idx = list(range(len(records))) | |
| rng.shuffle(idx) | |
| n_train = int(len(idx) * split_ratio) | |
| train = [records[i] for i in idx[:n_train]] | |
| val = [records[i] for i in idx[n_train:]] | |
| return train, val | |
| def _set_cosine_lrs( | |
| optimizer: AdamW, current_step: int, total_steps: int, min_lr_scale: float | |
| ) -> None: | |
| if total_steps <= 1: | |
| mult = 1.0 | |
| else: | |
| progress = min(1.0, max(0.0, current_step / float(total_steps - 1))) | |
| mult = min_lr_scale + (1.0 - min_lr_scale) * 0.5 * ( | |
| 1.0 + math.cos(math.pi * progress) | |
| ) | |
| for group in optimizer.param_groups: | |
| base_lr = group.get("initial_lr", group["lr"]) | |
| group["lr"] = base_lr * mult | |
| def evaluate( | |
| model: CompoundMidiTextContrastiveModel, | |
| loader: DataLoader, | |
| device: torch.device, | |
| ) -> Dict[str, float]: | |
| model.eval() | |
| loss_sum = 0.0 | |
| n_batches = 0 | |
| midi_chunks: List[torch.Tensor] = [] | |
| text_chunks: List[torch.Tensor] = [] | |
| for batch in loader: | |
| out = model( | |
| compound_input=batch["compound_input"].to(device), | |
| captions=batch["captions"], | |
| ) | |
| loss_sum += float(out["loss"].item()) | |
| midi_chunks.append(out["midi_embeds"].detach().cpu()) | |
| text_chunks.append(out["text_embeds"].detach().cpu()) | |
| n_batches += 1 | |
| model.train() | |
| if n_batches == 0: | |
| return {"loss": 0.0, "r1_m2t": 0.0, "r1_t2m": 0.0} | |
| midi_all = torch.cat(midi_chunks, dim=0) | |
| text_all = torch.cat(text_chunks, dim=0) | |
| logits = midi_all @ text_all.t() | |
| labels = torch.arange(logits.size(0)) | |
| r1_m2t = float((torch.argmax(logits, dim=1) == labels).float().mean().item()) | |
| r1_t2m = float((torch.argmax(logits.t(), dim=1) == labels).float().mean().item()) | |
| return {"loss": loss_sum / n_batches, "r1_m2t": r1_m2t, "r1_t2m": r1_t2m} | |
| def _load_compound_gpt(ckpt_path: Path, device: torch.device) -> CompoundGPT: | |
| ckpt = torch.load(ckpt_path, map_location=device, weights_only=True) | |
| cfg = default_compound_config() | |
| raw_cfg = ckpt.get("config") if isinstance(ckpt, dict) else None | |
| if isinstance(raw_cfg, dict): | |
| for k in CompoundGPTConfig.__dataclass_fields__.keys(): | |
| if k in raw_cfg: | |
| setattr(cfg, k, raw_cfg[k]) | |
| model = CompoundGPT(cfg).to(device) | |
| state = ckpt.get("model_state_dict", ckpt) | |
| model.load_state_dict(state, strict=False) | |
| model.eval() | |
| return model | |
| def train(args: argparse.Namespace) -> None: | |
| device = _pick_device() | |
| print(f"[compound-contrastive] device={device}") | |
| torch.manual_seed(args.seed) | |
| if torch.cuda.is_available(): | |
| torch.cuda.manual_seed_all(args.seed) | |
| all_records = _load_jsonl_records(Path(args.captions_jsonl)) | |
| train_records, val_records = _split_records( | |
| all_records, split_ratio=args.split_ratio, seed=args.seed | |
| ) | |
| print( | |
| "[compound-contrastive] records total/train/val=" | |
| f"{len(all_records)}/{len(train_records)}/{len(val_records)}" | |
| ) | |
| train_ds = CompoundCaptionDataset( | |
| records=train_records, | |
| is_train=True, | |
| max_seq_len=args.max_seq_len, | |
| seed=args.seed, | |
| ) | |
| val_ds = CompoundCaptionDataset( | |
| records=val_records, | |
| is_train=False, | |
| max_seq_len=args.max_seq_len, | |
| seed=args.seed + 1, | |
| ) | |
| train_loader = DataLoader( | |
| train_ds, | |
| batch_size=args.batch_size, | |
| shuffle=True, | |
| drop_last=True, | |
| num_workers=args.num_workers, | |
| pin_memory=torch.cuda.is_available(), | |
| collate_fn=_collate, | |
| ) | |
| val_loader = DataLoader( | |
| val_ds, | |
| batch_size=args.batch_size, | |
| shuffle=False, | |
| drop_last=False, | |
| num_workers=args.num_workers, | |
| pin_memory=torch.cuda.is_available(), | |
| collate_fn=_collate, | |
| ) | |
| midi_encoder = _load_compound_gpt( | |
| ckpt_path=Path(args.compound_checkpoint), | |
| device=device, | |
| ) | |
| model = CompoundMidiTextContrastiveModel( | |
| midi_compound_gpt=midi_encoder, | |
| text_model_name=args.text_model, | |
| embed_dim=args.embed_dim, | |
| init_temperature=args.init_temperature, | |
| min_temperature=args.min_temperature, | |
| max_temperature=args.max_temperature, | |
| device=device, | |
| ) | |
| model.train() | |
| optimizer = AdamW( | |
| [ | |
| { | |
| "params": model.midi_projection.parameters(), | |
| "lr": args.proj_lr, | |
| "initial_lr": args.proj_lr, | |
| "weight_decay": args.proj_weight_decay, | |
| }, | |
| { | |
| "params": model.text_projection.parameters(), | |
| "lr": args.proj_lr, | |
| "initial_lr": args.proj_lr, | |
| "weight_decay": args.proj_weight_decay, | |
| }, | |
| { | |
| "params": [model.log_temperature], | |
| "lr": args.temp_lr, | |
| "initial_lr": args.temp_lr, | |
| "weight_decay": 0.0, | |
| }, | |
| ] | |
| ) | |
| steps_per_epoch = len(train_loader) | |
| total_steps = max(1, args.epochs * steps_per_epoch) | |
| global_step = 0 | |
| best_val = float("inf") | |
| text_group_added = False | |
| start_epoch = 1 | |
| if args.resume_checkpoint: | |
| ckpt = torch.load( | |
| args.resume_checkpoint, | |
| map_location=device, | |
| weights_only=False, | |
| ) | |
| model.load_state_dict(ckpt["model_state_dict"]) | |
| if int(ckpt["epoch"]) >= args.unfreeze_text_epoch: | |
| model.unfreeze_text_encoder() | |
| optimizer.add_param_group( | |
| { | |
| "params": model.text_encoder.parameters(), | |
| "lr": args.text_lr, | |
| "initial_lr": args.text_lr, | |
| "weight_decay": args.text_weight_decay, | |
| } | |
| ) | |
| text_group_added = True | |
| optimizer.load_state_dict(ckpt["optimizer_state_dict"]) | |
| start_epoch = int(ckpt["epoch"]) + 1 | |
| global_step = int(ckpt["global_step"]) | |
| best_val = float(ckpt["val_loss"]) | |
| print( | |
| "[compound-contrastive] resumed from " | |
| f"epoch={ckpt['epoch']} val_loss={best_val:.4f}" | |
| ) | |
| results_dir = Path(args.results_dir) | |
| ckpt_dir = results_dir / "checkpoints_contrastive_compound" | |
| ckpt_dir.mkdir(parents=True, exist_ok=True) | |
| results_dir.mkdir(parents=True, exist_ok=True) | |
| log_csv = results_dir / "contrastive_compound_training_log.csv" | |
| if not log_csv.exists(): | |
| with open(log_csv, "w", newline="") as f: | |
| csv.DictWriter( | |
| f, | |
| fieldnames=[ | |
| "epoch", | |
| "step", | |
| "stage", | |
| "train_loss", | |
| "train_acc_m2t", | |
| "train_acc_t2m", | |
| "val_loss", | |
| "val_r1_m2t", | |
| "val_r1_t2m", | |
| "temperature", | |
| ], | |
| ).writeheader() | |
| t0 = time.perf_counter() | |
| for epoch in range(start_epoch, args.epochs + 1): | |
| if epoch == args.unfreeze_text_epoch and not text_group_added: | |
| model.unfreeze_text_encoder() | |
| optimizer.add_param_group( | |
| { | |
| "params": model.text_encoder.parameters(), | |
| "lr": args.text_lr, | |
| "initial_lr": args.text_lr, | |
| "weight_decay": args.text_weight_decay, | |
| } | |
| ) | |
| text_group_added = True | |
| print(f"[compound-contrastive] epoch={epoch}: unfroze text encoder.") | |
| stage = "A" if epoch < args.unfreeze_text_epoch else "B" | |
| train_loss_sum = 0.0 | |
| train_acc_m2t_sum = 0.0 | |
| train_acc_t2m_sum = 0.0 | |
| train_batches = 0 | |
| for batch in train_loader: | |
| _set_cosine_lrs( | |
| optimizer=optimizer, | |
| current_step=global_step, | |
| total_steps=total_steps, | |
| min_lr_scale=args.min_lr_scale, | |
| ) | |
| optimizer.zero_grad(set_to_none=True) | |
| out = model( | |
| compound_input=batch["compound_input"].to(device), | |
| captions=batch["captions"], | |
| ) | |
| loss = out["loss"] | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip_norm) | |
| optimizer.step() | |
| train_loss_sum += float(loss.item()) | |
| train_acc_m2t_sum += float(out["acc_midi_to_text"].item()) | |
| train_acc_t2m_sum += float(out["acc_text_to_midi"].item()) | |
| train_batches += 1 | |
| global_step += 1 | |
| train_loss = train_loss_sum / max(1, train_batches) | |
| train_acc_m2t = train_acc_m2t_sum / max(1, train_batches) | |
| train_acc_t2m = train_acc_t2m_sum / max(1, train_batches) | |
| val = evaluate(model, val_loader, device) | |
| temp_val = float(model.get_temperature().item()) | |
| print( | |
| f"[compound-contrastive] epoch={epoch}/{args.epochs} stage={stage} " | |
| f"train_loss={train_loss:.4f} val_loss={val['loss']:.4f} " | |
| f"train_acc=({train_acc_m2t:.3f},{train_acc_t2m:.3f}) " | |
| f"val_r1=({val['r1_m2t']:.3f},{val['r1_t2m']:.3f}) " | |
| f"temp={temp_val:.4f}" | |
| ) | |
| with open(log_csv, "a", newline="") as f: | |
| csv.DictWriter( | |
| f, | |
| fieldnames=[ | |
| "epoch", | |
| "step", | |
| "stage", | |
| "train_loss", | |
| "train_acc_m2t", | |
| "train_acc_t2m", | |
| "val_loss", | |
| "val_r1_m2t", | |
| "val_r1_t2m", | |
| "temperature", | |
| ], | |
| ).writerow( | |
| { | |
| "epoch": epoch, | |
| "step": global_step, | |
| "stage": stage, | |
| "train_loss": f"{train_loss:.6f}", | |
| "train_acc_m2t": f"{train_acc_m2t:.6f}", | |
| "train_acc_t2m": f"{train_acc_t2m:.6f}", | |
| "val_loss": f"{val['loss']:.6f}", | |
| "val_r1_m2t": f"{val['r1_m2t']:.6f}", | |
| "val_r1_t2m": f"{val['r1_t2m']:.6f}", | |
| "temperature": f"{temp_val:.6f}", | |
| } | |
| ) | |
| ckpt = { | |
| "model_state_dict": model.state_dict(), | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "epoch": epoch, | |
| "global_step": global_step, | |
| "val_loss": float(val["loss"]), | |
| "args": vars(args), | |
| } | |
| torch.save(ckpt, ckpt_dir / "clap_compound_latest.pt") | |
| if epoch % args.checkpoint_every == 0: | |
| torch.save(ckpt, ckpt_dir / f"clap_compound_epoch_{epoch:03d}.pt") | |
| if val["loss"] < best_val: | |
| best_val = val["loss"] | |
| torch.save(ckpt, ckpt_dir / "clap_compound_best.pt") | |
| elapsed = time.perf_counter() - t0 | |
| print( | |
| f"[compound-contrastive] finished in {elapsed/60:.1f} min, " | |
| f"best_val={best_val:.4f}" | |
| ) | |
| def parse_args() -> argparse.Namespace: | |
| p = argparse.ArgumentParser(description="Train compound MIDI-text contrastive model") | |
| p.add_argument( | |
| "--captions-jsonl", | |
| type=str, | |
| default=str(_ROOT / "data" / "midicaps_100.jsonl"), | |
| ) | |
| p.add_argument( | |
| "--compound-checkpoint", | |
| type=str, | |
| default=str( | |
| _ROOT / "results" / "test_compound" / "checkpoints_compound" / "compound_best.pt" | |
| ), | |
| ) | |
| p.add_argument( | |
| "--text-model", | |
| type=str, | |
| default="sentence-transformers/all-MiniLM-L6-v2", | |
| ) | |
| p.add_argument("--results-dir", type=str, default=str(_ROOT / "results" / "test_compound")) | |
| p.add_argument("--epochs", type=int, default=10) | |
| p.add_argument("--unfreeze-text-epoch", type=int, default=6) | |
| p.add_argument("--batch-size", type=int, default=16) | |
| p.add_argument("--max-seq-len", type=int, default=512) | |
| p.add_argument("--split-ratio", type=float, default=0.95) | |
| p.add_argument("--num-workers", type=int, default=4) | |
| p.add_argument("--seed", type=int, default=17) | |
| p.add_argument("--embed-dim", type=int, default=256) | |
| p.add_argument("--init-temperature", type=float, default=0.07) | |
| p.add_argument("--min-temperature", type=float, default=0.01) | |
| p.add_argument("--max-temperature", type=float, default=1.0) | |
| p.add_argument("--proj-lr", type=float, default=1e-4) | |
| p.add_argument("--temp-lr", type=float, default=1e-5) | |
| p.add_argument("--text-lr", type=float, default=1e-5) | |
| p.add_argument("--proj-weight-decay", type=float, default=0.01) | |
| p.add_argument("--text-weight-decay", type=float, default=0.01) | |
| p.add_argument("--min-lr-scale", type=float, default=0.01) | |
| p.add_argument("--grad-clip-norm", type=float, default=1.0) | |
| p.add_argument("--checkpoint-every", type=int, default=5) | |
| p.add_argument("--resume-checkpoint", type=str, default="") | |
| return p.parse_args() | |
| if __name__ == "__main__": | |
| train(parse_args()) | |