Unclear usage instructions
#1
by nshmyrevgmail - opened
Couple issues:
- Current readme doesn't describe how to create the model object
- 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)