saga / modeling_saga.py
AndreasEefsen's picture
Upload folder using huggingface_hub
5056f92 verified
Raw
History Blame Contribute Delete
2.35 kB
import torch
import numpy as np
from transformers import PreTrainedModel
class SagaForCausalLM(PreTrainedModel):
config_class = None
def __init__(self, config):
super().__init__(config)
@torch.inference_mode()
def transcribe(self, audio, processor):
step_seconds = 15.0
step_samples = max(1, int(round(step_seconds * processor.target_sr)))
base_prompt = processor.get_prompt()
raw_decoded = ""
audio_accum = np.zeros((0,), dtype=np.float32)
for chunk_index, start in enumerate(range(0, audio.shape[0], step_samples)):
chunk = audio[start : start + step_samples]
if chunk.size == 0:
continue
if audio_accum.size == 0:
audio_accum = chunk
else:
audio_accum = np.concatenate([audio_accum, chunk], axis=0)
prefix = ""
if chunk_index >= 1 and raw_decoded:
cur_ids = processor.tokenizer.encode(raw_decoded)
rollback = 8
while True:
end_idx = max(0, len(cur_ids) - rollback)
prefix = processor.tokenizer.decode(cur_ids[:end_idx]) if end_idx > 0 else ""
if "\ufffd" not in prefix or end_idx == 0:
break
rollback += 1
prompt = base_prompt + prefix
inputs = processor(
text=[prompt],
audio=[audio_accum],
sampling_rate=processor.target_sr,
return_tensors="pt",
padding=True,
)
inputs = {key: value.to(self.device) for key, value in inputs.items()}
if "input_features" in inputs and inputs["input_features"].is_floating_point():
inputs["input_features"] = inputs["input_features"].to(dtype=self.dtype)
generated = self.generate(
**inputs,
max_new_tokens=2048,
)
decoded = processor.batch_decode(
generated.sequences[:, inputs["input_ids"].shape[1]:],
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)[0]
raw_decoded = prefix + decoded
return raw_decoded.strip()