Files
2026-07-19 23:52:25 +04:00

323 lines
9.9 KiB
Python

import os
import sys
os.environ.setdefault("HF_HOME", "/mnt/D/.cache/huggingface")
os.environ.setdefault("HF_DATASETS_CACHE", "/mnt/D/.cache/huggingface/datasets")
os.environ.setdefault("TMPDIR", "/mnt/D/tmp")
os.makedirs(os.environ["TMPDIR"], exist_ok=True)
os.makedirs(os.environ["HF_DATASETS_CACHE"], exist_ok=True)
print("=== initializing environment")
from dataclasses import dataclass
from typing import Any
import jiwer
from datasets import load_dataset, Audio
from peft import LoraConfig, TaskType, get_peft_model
import numpy as np
import torch
from tqdm import tqdm as _tqdm
from transformers import (
Seq2SeqTrainer,
Seq2SeqTrainingArguments,
WhisperForConditionalGeneration,
WhisperProcessor,
)
class WhisperLoraTrainer(Seq2SeqTrainer):
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
input_features = inputs.get("input_features")
labels = inputs.get("labels")
dtype = next(model.parameters()).dtype
input_features = input_features.to(dtype)
labels = labels.to(torch.long)
outputs = model.base_model(input_features=input_features, labels=labels)
loss = outputs.loss
return (loss, outputs) if return_outputs else loss
def prediction_step(self, model, inputs, prediction_loss_only, ignore_keys=None):
if not self.args.predict_with_generate or prediction_loss_only:
return super().prediction_step(
model, inputs, prediction_loss_only, ignore_keys=ignore_keys
)
inputs = self._prepare_inputs(inputs)
labels = inputs.get("labels")
dtype = next(model.parameters()).dtype
with torch.no_grad():
generated_tokens = model.base_model.model.generate(
input_features=inputs["input_features"].to(dtype),
language=LANGUAGE,
task="transcribe",
)
loss = None
if labels is not None:
with torch.no_grad():
outputs = model.base_model(
input_features=inputs["input_features"].to(dtype),
labels=labels.to(torch.long),
)
loss = outputs.loss.detach().float()
generated_tokens = generated_tokens.detach().cpu().to(torch.long)
if labels is not None:
labels = labels.detach().cpu().to(torch.long)
return loss, generated_tokens, labels
print("[+] imports ok")
MODEL_ID = "openai/whisper-medium"
DATASET_ID = "fsicoli/common_voice_17_0"
LANGUAGE = "ru"
SAMPLING_RATE = 16000
OUTPUT_DIR = "./whisper-medium-ru-lora"
EVAL_SAMPLES = 1000
LORA_R = 32
LORA_ALPHA = 64
LORA_DROPOUT = 0.05
@dataclass
class DataCollatorSpeechSeq2SeqWithPadding:
processor: Any
decoder_start_token_id: int
def __call__(self, features: list[dict]) -> dict:
input_features = [{"input_features": f["input_features"]} for f in features]
batch = self.processor.feature_extractor.pad(input_features, return_tensors="pt")
label_features = [{"input_ids": f["labels"]} for f in features]
labels_batch = self.processor.tokenizer.pad(label_features, return_tensors="pt")
labels = labels_batch["input_ids"].masked_fill(
labels_batch.attention_mask.ne(1), -100
)
if (labels[:, 0] == self.decoder_start_token_id).all().cpu().item():
labels = labels[:, 1:]
batch["labels"] = labels
return batch
def load_processor_and_model():
print("=== loading processor")
try:
processor = WhisperProcessor.from_pretrained(
MODEL_ID, language=LANGUAGE, task="transcribe"
)
print("[+] processor loaded")
except Exception as e:
print(f"[-] processor load failed: {e}")
sys.exit(1)
print("=== loading model")
try:
model = WhisperForConditionalGeneration.from_pretrained(MODEL_ID)
model.config.use_cache = False
model.config.forced_decoder_ids = None
model.config.suppress_tokens = []
model.generation_config.language = LANGUAGE
model.generation_config.task = "transcribe"
model.generation_config.forced_decoder_ids = None
model.generation_config.max_length = None
print("[+] model loaded")
except Exception as e:
print(f"[-] model load failed: {e}")
sys.exit(1)
return processor, model
def apply_lora(model):
print("=== applying LoRA")
try:
lora_config = LoraConfig(
r=LORA_R,
lora_alpha=LORA_ALPHA,
lora_dropout=LORA_DROPOUT,
target_modules=["q_proj", "k_proj", "v_proj", "out_proj"],
task_type=TaskType.SEQ_2_SEQ_LM,
)
model = get_peft_model(model, lora_config)
model.enable_input_require_grads()
model.print_trainable_parameters()
print("[+] LoRA applied")
except Exception as e:
print(f"[-] LoRA apply failed: {e}")
sys.exit(1)
return model
def prepare_batch(batch, processor):
audio = batch["audio"]
batch["input_features"] = processor.feature_extractor(
audio["array"], sampling_rate=audio["sampling_rate"]
).input_features[0]
batch["labels"] = processor.tokenizer(batch["sentence"]).input_ids
return batch
def load_and_prepare_dataset(processor):
print("=== loading dataset (train + validation only)")
try:
raw_train, raw_eval = load_dataset(
DATASET_ID,
LANGUAGE,
split=["train", "validation"],
trust_remote_code=True,
)
print("[+] dataset loaded")
except Exception as e:
print(f"[-] dataset load failed: {e}")
sys.exit(1)
print(f"=== slicing validation to {EVAL_SAMPLES} samples")
raw_eval = raw_eval.select(range(EVAL_SAMPLES))
print(f"[+] validation sliced to {len(raw_eval)} samples")
print("=== casting audio column")
try:
raw_train = raw_train.cast_column("audio", Audio(sampling_rate=SAMPLING_RATE))
raw_eval = raw_eval.cast_column("audio", Audio(sampling_rate=SAMPLING_RATE))
print("[+] audio cast done")
except Exception as e:
print(f"[-] audio cast failed: {e}")
sys.exit(1)
fn = lambda batch: prepare_batch(batch, processor)
print("=== mapping train split")
try:
train_dataset = raw_train.map(
fn,
remove_columns=raw_train.column_names,
num_proc=1,
cache_file_name=str(os.path.join(os.environ["HF_DATASETS_CACHE"], "cv17_ru_train.arrow")),
)
print("[+] train split mapped")
except Exception as e:
print(f"[-] train map failed: {e}")
sys.exit(1)
print("=== mapping validation split")
try:
eval_dataset = raw_eval.map(
fn,
remove_columns=raw_eval.column_names,
num_proc=1,
load_from_cache_file=False,
)
print("[+] validation split mapped")
except Exception as e:
print(f"[-] validation map failed: {e}")
sys.exit(1)
return train_dataset, eval_dataset
def build_compute_metrics(processor):
def compute_metrics(pred):
pred_ids = pred.predictions
label_ids = pred.label_ids
if isinstance(pred_ids, tuple):
pred_ids = pred_ids[0]
label_ids = np.array(label_ids, dtype=np.int64)
label_ids[label_ids == -100] = processor.tokenizer.pad_token_id
pred_str = [
processor.tokenizer.decode(
[int(t) for t in seq if 0 <= int(t) < processor.tokenizer.vocab_size],
skip_special_tokens=True,
)
for seq in pred_ids
]
label_str = processor.tokenizer.batch_decode(label_ids, skip_special_tokens=True)
wer = 100 * jiwer.wer(label_str, pred_str)
_tqdm.write(f"[+] eval wer: {wer:.2f}%")
return {"wer": wer}
return compute_metrics
def main():
processor, model = load_processor_and_model()
model = apply_lora(model)
train_dataset, eval_dataset = load_and_prepare_dataset(processor)
data_collator = DataCollatorSpeechSeq2SeqWithPadding(
processor=processor,
decoder_start_token_id=model.config.decoder_start_token_id,
)
print("[+] data collator configured")
compute_metrics = build_compute_metrics(processor)
print("=== configuring training arguments")
training_args = Seq2SeqTrainingArguments(
output_dir=OUTPUT_DIR,
per_device_train_batch_size=2,
gradient_accumulation_steps=1,
learning_rate=5e-5,
warmup_steps=500,
max_steps=4000,
gradient_checkpointing=True,
predict_with_generate=True,
per_device_eval_batch_size=2,
generation_max_length=64,
save_steps=1000,
logging_steps=100,
fp16=True,
fp16_full_eval=False,
eval_strategy="steps",
eval_steps=1000,
metric_for_best_model="wer",
greater_is_better=False,
load_best_model_at_end=True,
generation_config=None,
report_to=["tensorboard"],
)
print("[+] training arguments configured")
print("=== initializing trainer")
try:
trainer = WhisperLoraTrainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
data_collator=data_collator,
compute_metrics=compute_metrics,
processing_class=processor.feature_extractor,
)
print("[+] trainer initialized")
except Exception as e:
print(f"[-] trainer init failed: {e}")
sys.exit(1)
print("=== starting training")
trainer.train()
print("[+] training complete")
print("=== saving model")
try:
model.save_pretrained(OUTPUT_DIR)
processor.save_pretrained(OUTPUT_DIR)
print(f"[+] model saved to {OUTPUT_DIR}")
except Exception as e:
print(f"[-] save failed: {e}")
sys.exit(1)
if __name__ == "__main__":
main()