onepunchgin commited on
Commit
b15d0f9
·
verified ·
1 Parent(s): b318d88

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +7 -8
app.py CHANGED
@@ -3,23 +3,22 @@ import torchaudio
3
  import torch
4
  from transformers import Wav2Vec2Processor, Wav2Vec2ForCTC
5
 
6
- # Load local model
7
- model_dir = "ccc_wav2vec_model"
8
- processor = Wav2Vec2Processor.from_pretrained(model_dir)
9
- model = Wav2Vec2ForCTC.from_pretrained(model_dir)
 
10
  model.eval()
11
 
12
  def transcribe(audio):
13
  waveform, sr = torchaudio.load(audio)
14
  if sr != 16000:
15
- resampler = torchaudio.transforms.Resample(orig_freq=sr, new_freq=16000)
16
- waveform = resampler(waveform)
17
  inputs = processor(waveform.squeeze().numpy(), sampling_rate=16000, return_tensors="pt", padding=True)
18
  with torch.no_grad():
19
  logits = model(**inputs).logits
20
  predicted_ids = torch.argmax(logits, dim=-1)
21
- transcription = processor.batch_decode(predicted_ids)[0]
22
- return transcription
23
 
24
  iface = gr.Interface(fn=transcribe,
25
  inputs=gr.Audio(type="filepath"),
 
3
  import torch
4
  from transformers import Wav2Vec2Processor, Wav2Vec2ForCTC
5
 
6
+ # Load model directly from Hugging Face Hub
7
+ model_id = "onepunchgin/ASR-Model"
8
+
9
+ processor = Wav2Vec2Processor.from_pretrained(model_id)
10
+ model = Wav2Vec2ForCTC.from_pretrained(model_id)
11
  model.eval()
12
 
13
  def transcribe(audio):
14
  waveform, sr = torchaudio.load(audio)
15
  if sr != 16000:
16
+ waveform = torchaudio.transforms.Resample(orig_freq=sr, new_freq=16000)(waveform)
 
17
  inputs = processor(waveform.squeeze().numpy(), sampling_rate=16000, return_tensors="pt", padding=True)
18
  with torch.no_grad():
19
  logits = model(**inputs).logits
20
  predicted_ids = torch.argmax(logits, dim=-1)
21
+ return processor.batch_decode(predicted_ids)[0]
 
22
 
23
  iface = gr.Interface(fn=transcribe,
24
  inputs=gr.Audio(type="filepath"),