AudioModel-v1 / build_data.py
TobiasLogic's picture
Upload build_data.py with huggingface_hub
f48ab3f verified
Raw
History Blame Contribute Delete
5.35 kB
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()