| from __future__ import annotations |
|
|
| import argparse |
| import csv |
| import json |
| import os |
| import time |
| from multiprocessing import Pool |
|
|
| import numpy as np |
| import torch |
| from diffusers.pipelines.deprecated.audio_diffusion.mel import Mel |
| from transformers import AutoTokenizer, ClapModel |
|
|
| X_RES, Y_RES = 384, 256 |
| HOP, NFFT, SR, TOPDB = 1024, 2048, 22050, 80 |
| MAX_TOKENS = 32 |
|
|
| def read_captions(csv_path): |
| rows = [] |
| with open(csv_path, newline="", encoding="utf-8", errors="replace") as f: |
| reader = csv.DictReader(f) |
| for row in reader: |
| fn = row["file_name"] |
| caps = [row[f"caption_{i}"] for i in range(1, 6) if row.get(f"caption_{i}")] |
| rows.append((fn, caps)) |
| return rows |
|
|
| _mel = None |
| def _get_mel(): |
| global _mel |
| if _mel is None: |
| _mel = Mel(x_res=X_RES, y_res=Y_RES, sample_rate=SR, n_fft=NFFT, hop_length=HOP, top_db=TOPDB) |
| return _mel |
|
|
| def to_mel_array(path): |
| try: |
| mel = _get_mel() |
| mel.load_audio(audio_file=path) |
| img = mel.audio_slice_to_image(0) |
| return path, np.array(img, dtype=np.uint8) |
| except Exception: |
| return path, None |
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--audio-dir", required=True) |
| ap.add_argument("--captions-csv", required=True) |
| ap.add_argument("--split", required=True) |
| ap.add_argument("--out", default="/root/data") |
| ap.add_argument("--clap", default="laion/clap-htsat-unfused") |
| ap.add_argument("--workers", type=int, default=16) |
| ap.add_argument("--batch", type=int, default=128) |
| ap.add_argument("--device", default="cuda") |
| args = ap.parse_args() |
| os.makedirs(args.out, exist_ok=True) |
|
|
| rows = read_captions(args.captions_csv) |
| print(f"[build] {args.split}: {len(rows)} clips listed in captions csv", flush=True) |
|
|
| paths = [os.path.join(args.audio_dir, fn) for fn, _ in rows] |
| missing = [p for p in paths if not os.path.exists(p)] |
| if missing: |
| print(f"[build] WARNING {len(missing)} missing audio files, e.g. {missing[:3]}", flush=True) |
|
|
| t0 = time.time() |
| mel_by_path, failed = {}, 0 |
| with Pool(args.workers) as pool: |
| for i, (p, arr) in enumerate(pool.imap(to_mel_array, paths, chunksize=8)): |
| if arr is None: |
| failed += 1 |
| else: |
| mel_by_path[p] = arr |
| if (i + 1) % 500 == 0: |
| print(f"[build] mel {i+1}/{len(paths)} failed={failed} {time.time()-t0:.0f}s", flush=True) |
| print(f"[build] mel done: {len(mel_by_path)} ok, {failed} failed, {time.time()-t0:.0f}s", flush=True) |
|
|
| clip_paths = list(mel_by_path.keys()) |
| clip_index = {p: i for i, p in enumerate(clip_paths)} |
| mel_arr = np.stack([mel_by_path[p] for p in clip_paths]) |
| print(f"[build] mel_arr {mel_arr.shape} {mel_arr.nbytes/2**20:.0f} MiB", flush=True) |
|
|
| pair_clip_idx, pair_captions = [], [] |
| for fn, caps in rows: |
| p = os.path.join(args.audio_dir, fn) |
| if p not in clip_index: |
| continue |
| for c in caps: |
| pair_clip_idx.append(clip_index[p]) |
| pair_captions.append(c) |
| print(f"[build] {len(pair_captions)} (clip, caption) pairs", flush=True) |
|
|
| dev = args.device |
| tok = AutoTokenizer.from_pretrained(args.clap) |
| clap = ClapModel.from_pretrained(args.clap).to(dev).eval() |
| pool_dim = clap.config.projection_dim |
| seq_dim = clap.config.text_config.hidden_size |
|
|
| @torch.no_grad() |
| def encode(strings): |
| enc = tok(strings, padding="max_length", truncation=True, max_length=MAX_TOKENS, |
| return_tensors="pt").to(dev) |
| out = clap.text_model(**enc) |
| seq = out.last_hidden_state.float() |
| pooled = clap.text_projection(out.pooler_output).float() |
| return seq.cpu().numpy().astype(np.float16), pooled.cpu().numpy().astype(np.float16) |
|
|
| seq_chunks, pool_chunks = [], [] |
| t0 = time.time() |
| for i in range(0, len(pair_captions), args.batch): |
| chunk = pair_captions[i:i + args.batch] |
| s, p = encode(chunk) |
| seq_chunks.append(s) |
| pool_chunks.append(p) |
| if (i // args.batch) % 20 == 0: |
| print(f"[build] clap {i}/{len(pair_captions)} {time.time()-t0:.0f}s", flush=True) |
| text_seq = np.concatenate(seq_chunks) |
| text_pool = np.concatenate(pool_chunks) |
| print(f"[build] text_seq {text_seq.shape} text_pool {text_pool.shape}", flush=True) |
|
|
| np.save(f"{args.out}/{args.split}_mel.npy", mel_arr) |
| np.save(f"{args.out}/{args.split}_text_seq.npy", text_seq) |
| np.save(f"{args.out}/{args.split}_text_pool.npy", text_pool) |
| json.dump({"pair_clip_idx": pair_clip_idx, "captions": pair_captions, |
| "clip_files": [os.path.basename(p) for p in clip_paths], |
| "n_clips": len(clip_paths), "seq_dim": seq_dim, "pool_dim": pool_dim, |
| "x_res": X_RES, "y_res": Y_RES, "hop_length": HOP, "n_fft": NFFT, |
| "sample_rate": SR, "top_db": TOPDB, "max_tokens": MAX_TOKENS}, |
| open(f"{args.out}/{args.split}_meta.json", "w")) |
|
|
| if args.split == "development": |
| ns, npz = encode([""]) |
| np.save(f"{args.out}/null_seq.npy", ns[0]) |
| np.save(f"{args.out}/null_pool.npy", npz[0]) |
| print("[build] wrote null embeddings", flush=True) |
|
|
| print("BUILDDONE", flush=True) |
|
|
| if __name__ == "__main__": |
| main() |
|
|