init
This commit is contained in:
+175
@@ -0,0 +1,175 @@
|
||||
import sys
|
||||
import torch
|
||||
import datasets
|
||||
from transformers import (
|
||||
AutoProcessor,
|
||||
WhisperForConditionalGeneration,
|
||||
Seq2SeqTrainer,
|
||||
Seq2SeqTrainingArguments
|
||||
)
|
||||
from peft import LoraConfig, get_peft_model
|
||||
from jiwer import wer
|
||||
|
||||
datasets.disable_caching()
|
||||
|
||||
# === Configuration ===
|
||||
DATASET = "fsicoli/common_voice_17_0"
|
||||
LANG = "ru"
|
||||
SAMPLING_RATE = 16000
|
||||
MODEL_NAME = "openai/whisper-small"
|
||||
TRAIN_BATCH = 2
|
||||
EVAL_BATCH = 2
|
||||
MAX_LENGTH = 64
|
||||
LR = 1e-4
|
||||
MAX_STEPS = 4000
|
||||
WARMUP_STEPS = 100
|
||||
BATCH_SIZE = 512
|
||||
|
||||
print("=== Loading processor", flush=True)
|
||||
processor = AutoProcessor.from_pretrained(MODEL_NAME, language=LANG, task="transcribe")
|
||||
print("[+] Processor loaded", flush=True)
|
||||
|
||||
print("=== Loading model", flush=True)
|
||||
model = WhisperForConditionalGeneration.from_pretrained(MODEL_NAME)
|
||||
print("[+] Model loaded", flush=True)
|
||||
|
||||
print("=== Configuring LoRA", flush=True)
|
||||
lora_config = LoraConfig(
|
||||
task_type="SEQ_2_SEQ_LM",
|
||||
target_modules=["q_proj", "k_proj", "v_proj", "out_proj"],
|
||||
r=8,
|
||||
lora_alpha=32,
|
||||
lora_dropout=0.1,
|
||||
bias="none",
|
||||
fan_in_fan_out=False,
|
||||
)
|
||||
model = get_peft_model(model, lora_config)
|
||||
print("[+] LoRA configured", flush=True)
|
||||
|
||||
# === Streaming train dataset ===
|
||||
print("=== Loading train dataset", flush=True)
|
||||
train_ds = datasets.load_dataset(DATASET, LANG, split="train")
|
||||
train_ds = train_ds.remove_columns(["accent", "age", "client_id", "down_votes", "gender", "locale", "path", "segment", "up_votes"])
|
||||
print("[+] Train dataset loaded", flush=True)
|
||||
|
||||
# === Small cached validation/test subsets ===
|
||||
print("=== Loading validation/test datasets", flush=True)
|
||||
val_ds = datasets.load_dataset(DATASET, LANG, split="validation[:500]")
|
||||
val_ds = val_ds.remove_columns(["accent", "age", "client_id", "down_votes", "gender", "locale", "path", "segment", "up_votes"])
|
||||
test_ds = datasets.load_dataset(DATASET, LANG, split="test[:500]")
|
||||
test_ds = test_ds.remove_columns(["accent", "age", "client_id", "down_votes", "gender", "locale", "path", "segment", "up_votes"])
|
||||
print("[+] Validation/test datasets are loaded", flush=True)
|
||||
|
||||
# Cast audio column
|
||||
print("=== Casting datasets to 16kHz", flush=True)
|
||||
train_ds = train_ds.cast_column("audio", datasets.Audio(sampling_rate=SAMPLING_RATE))
|
||||
val_ds = val_ds.cast_column("audio", datasets.Audio(sampling_rate=SAMPLING_RATE))
|
||||
test_ds = test_ds.cast_column("audio", datasets.Audio(sampling_rate=SAMPLING_RATE))
|
||||
print("[+] Audio column cast to 16kHz", flush=True)
|
||||
|
||||
# === Map function with picklable args ===
|
||||
def prepare_examples(batch, processor_name=MODEL_NAME):
|
||||
texts = batch["sentence"]
|
||||
audios = [x["array"] for x in batch["audio"]]
|
||||
inputs = processor(
|
||||
audios,
|
||||
sampling_rate=SAMPLING_RATE,
|
||||
return_tensors="np"
|
||||
)
|
||||
batch["input_features"] = inputs.input_features
|
||||
|
||||
batch["labels"] = processor.tokenizer(
|
||||
texts,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=MAX_LENGTH
|
||||
).input_ids
|
||||
return batch
|
||||
|
||||
# === Preprocessing ===
|
||||
try:
|
||||
print("=== Preprocessing datasets", flush=True)
|
||||
train_ds = train_ds.map(
|
||||
prepare_examples,
|
||||
batched=True,
|
||||
batch_size=BATCH_SIZE,
|
||||
remove_columns=["audio", "sentence"],
|
||||
load_from_cache_file=False,
|
||||
)
|
||||
val_ds = val_ds.map(
|
||||
prepare_examples,
|
||||
batched=True,
|
||||
batch_size=BATCH_SIZE,
|
||||
remove_columns=["audio", "sentence"],
|
||||
load_from_cache_file=False,
|
||||
)
|
||||
test_ds = test_ds.map(
|
||||
prepare_examples,
|
||||
batched=True,
|
||||
batch_size=BATCH_SIZE,
|
||||
remove_columns=["audio", "sentence"],
|
||||
load_from_cache_file=False,
|
||||
)
|
||||
print("[+] Datasets are preprocessed", flush=True)
|
||||
except Exception as e:
|
||||
print(f"[-] Preprocessing failed: {e}", flush=True)
|
||||
sys.exit(1)
|
||||
|
||||
# === Data collator ===
|
||||
def data_collator(batch):
|
||||
input_features = torch.tensor([ex["input_features"] for ex in batch], dtype=torch.float32)
|
||||
labels = torch.tensor([ex["labels"] for ex in batch], dtype=torch.long)
|
||||
return {"input_features": input_features, "labels": labels}
|
||||
|
||||
# === Metrics ===
|
||||
def compute_metrics(pred):
|
||||
pred_ids = pred.predictions
|
||||
label_ids = pred.label_ids
|
||||
pred_str = processor.batch_decode(pred_ids, skip_special_tokens=True)
|
||||
label_str = processor.batch_decode(label_ids, skip_special_tokens=True)
|
||||
return {"wer": wer(label_str, pred_str)}
|
||||
|
||||
# === Training arguments ===
|
||||
training_args = Seq2SeqTrainingArguments(
|
||||
output_dir="./whisper_lora",
|
||||
per_device_train_batch_size=TRAIN_BATCH,
|
||||
gradient_accumulation_steps=1,
|
||||
learning_rate=LR,
|
||||
warmup_steps=WARMUP_STEPS,
|
||||
max_steps=MAX_STEPS,
|
||||
gradient_checkpointing=True,
|
||||
predict_with_generate=True,
|
||||
per_device_eval_batch_size=EVAL_BATCH,
|
||||
generation_max_length=MAX_LENGTH,
|
||||
save_steps=1000,
|
||||
logging_steps=100,
|
||||
fp16=True,
|
||||
eval_strategy="steps",
|
||||
eval_steps=1000,
|
||||
metric_for_best_model="wer",
|
||||
greater_is_better=False,
|
||||
load_best_model_at_end=True,
|
||||
report_to=["tensorboard"],
|
||||
)
|
||||
|
||||
# === Trainer ===
|
||||
print("=== Initializing trainer", flush=True)
|
||||
trainer = Seq2SeqTrainer(
|
||||
model=model,
|
||||
args=training_args,
|
||||
train_dataset=train_ds,
|
||||
eval_dataset=val_ds,
|
||||
data_collator=data_collator,
|
||||
compute_metrics=compute_metrics,
|
||||
)
|
||||
print("[+] Trainer initialized", flush=True)
|
||||
|
||||
# === Start training ===
|
||||
print("=== Starting training", flush=True)
|
||||
trainer.train()
|
||||
print("[+] Training finished", flush=True)
|
||||
|
||||
# === Final evaluation on test set ===
|
||||
print("=== Evaluating on test set", flush=True)
|
||||
results = trainer.evaluate(test_ds)
|
||||
print(f"[+] Test WER: {results['eval_wer']:.4f}", flush=True)
|
||||
Reference in New Issue
Block a user