Team Ai
Modelpublic

Mayank022/Audio-Language-Model

sourceHugging Faceapache-2.0updated 8mo agoView on Hugging Face
0likes
inference.py55 linesDownload Raw Back to root
1 2import torch3import torchaudio4import transformers5from config import ModelConfig6from model import MultiModalModel7 8def run_inference(audio_path: str, model_path: str = None):9    # Load Config & Model10    config = ModelConfig()11    12    13    model = MultiModalModel(config)14    15    if model_path:16        state_dict = torch.load(f"{model_path}/pytorch_model.bin", map_location="cpu")17        model.load_state_dict(state_dict, strict=False)18    19    model.eval()20    21    # Process Audio22    processor = transformers.AutoProcessor.from_pretrained(config.audio_model_id)23    audio, sr = torchaudio.load(audio_path)24    if sr != 16000:25        audio = torchaudio.functional.resample(audio, sr, 16000)26    if audio.shape[0] > 1:27        audio = audio.mean(dim=0, keepdim=True)28        29    audio_inputs = processor(audio.squeeze().numpy(), sampling_rate=16000, return_tensors="pt")30    audio_values = audio_inputs.input_features31    32    # Create Input Text33    tokenizer = transformers.AutoTokenizer.from_pretrained(config.text_model_id)34    text = "Transcribe the following audio:"35    text_inputs = tokenizer(text, return_tensors="pt")36    37    # Generate38    with torch.no_grad():39        generated_ids = model.generate(40            input_ids=text_inputs.input_ids,41            audio_values=audio_values,42            max_new_tokens=20043        )44    45    transcription = tokenizer.decode(generated_ids[0], skip_special_tokens=True)46    print("Transcription:", transcription)47    return transcription48 49if __name__ == "__main__":50    import sys51    if len(sys.argv) > 1:52        run_inference(sys.argv[1])53    else:54        print("Usage: python -m audio_lm.inference path/to/audio.wav")55