Unclear usage instructions

#1
by nshmyrevgmail - opened

Couple issues:

  1. Current readme doesn't describe how to create the model object
  2. Described way to compute spectrogram with 64-dim input is not compatible with the model, one better use embedded extractor (model.encoder.spec_frontend):

Something like this could work instead:

from asv.models.assemble import build_model

# Model configuration
model_config = {
    "encoder": {
        "_target_": "asv.models.redimnet2.ReDimNet2Encoder",
        "model_name": "b6",
        "pretrained": False,
        "train_type": "lm",
        "dataset": "vox2",
        "strict_load": False,
        "model_overrides": {
            "out_channels": 224,
            "return_2d_output": True,
        }
    },
    "bridge": {
        "_target_": "asv.models.bridge.IdentityBridge"
    },
    "classifier": {
        "_target_": "asv.models.classifier.IdentityClassifier"
    }
}


state_dict = torch.load("redimnet2-plus-model/pytorch_model_fsdp.bin", map_location="cuda")
# Strip FSDP wrapper prefixes
state_dict = {
    k.replace("_orig_mod.", "").replace("module.", ""): v
    for k, v in state_dict.items()
}

model = build_model(model_config)
model.load_state_dict(state_dict, strict=False)
model = model.to('cuda')
model.eval()
spec_frontend = model.encoder.spec_frontend

def load_audio(path, target_sr=16000):
    """Load audio file."""
    audio, sr = torchaudio.load(path)
    if audio.shape[0] > 1:
        audio = audio.mean(dim=0)  # Convert to mono
    if sr != target_sr:
        audio = resample(audio, sr, target_sr)
    return audio

def compute_embedding(audio_path, device="cuda"):
    audio = load_audio(audio_path)
    waveform = audio.unsqueeze(0).to(device)

    with torch.no_grad():
        mel = spec_frontend(waveform)
        embedding = model(mel)
        embedding = torch.nn.functional.normalize(embedding, dim=1)

    return embedding.cpu().squeeze(0)

Sign up or log in to comment