Diplomovka/strojovy_preklad/translate_dataset_to_slovak.py
2026-08-19 01:28:10 +02:00

174 lines
4.4 KiB
Python

import json
from pathlib import Path
import torch
from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
# -----------------------------
# Nastavenia
# -----------------------------
DATASET_NAME = "tatsu-lab/alpaca"
SPLIT = "train"
OUTPUT_DIR = Path("/home/schwarc/diplomovka/translated_datasets")
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
OUTPUT_FILE = OUTPUT_DIR / "alpaca_translated_en_sk_1000.jsonl"
MAX_SAMPLES = 1000
TRANSLATION_MODEL = "facebook/nllb-200-distilled-600M"
SOURCE_LANG = "eng_Latn"
TARGET_LANG = "slk_Latn"
BATCH_SIZE = 8
MAX_INPUT_LENGTH = 1024
MAX_NEW_TOKENS = 512
# Pomocné funkcie
def is_empty(value):
if value is None:
return True
value = str(value).strip()
return value == "" or value.lower() == "nan"
def load_already_done(path):
"""
Ak skript spadne alebo ho zastavíš, vie pokračovať.
Načíta už preložené riadky.
"""
if not path.exists():
return []
rows = []
with open(path, "r", encoding="utf-8") as f:
for line in f:
if line.strip():
rows.append(json.loads(line))
return rows
def translate_batch(texts, tokenizer, model, device):
"""
Preloží batch textov z angličtiny do slovenčiny.
Prázdne texty nechá prázdne.
"""
results = [""] * len(texts)
non_empty_indices = []
non_empty_texts = []
for i, text in enumerate(texts):
if is_empty(text):
results[i] = ""
else:
non_empty_indices.append(i)
non_empty_texts.append(str(text).strip())
if not non_empty_texts:
return results
tokenizer.src_lang = SOURCE_LANG
inputs = tokenizer(
non_empty_texts,
return_tensors="pt",
padding=True,
truncation=True,
max_length=MAX_INPUT_LENGTH,
).to(device)
forced_bos_token_id = tokenizer.convert_tokens_to_ids(TARGET_LANG)
with torch.no_grad():
generated_tokens = model.generate(
**inputs,
forced_bos_token_id=forced_bos_token_id,
max_new_tokens=MAX_NEW_TOKENS,
num_beams=4,
)
translated = tokenizer.batch_decode(
generated_tokens,
skip_special_tokens=True,
)
for idx, translation in zip(non_empty_indices, translated):
results[idx] = translation.strip()
return results
# Main
def main():
print("Loading dataset...")
dataset = load_dataset(DATASET_NAME, split=SPLIT)
if MAX_SAMPLES is not None:
dataset = dataset.select(range(min(MAX_SAMPLES, len(dataset))))
print(f"Dataset size: {len(dataset)}")
print("Loading translation model...")
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Device: {device}")
tokenizer = AutoTokenizer.from_pretrained(TRANSLATION_MODEL)
model = AutoModelForSeq2SeqLM.from_pretrained(
TRANSLATION_MODEL,
torch_dtype=torch.float16 if device == "cuda" else torch.float32,
).to(device)
model.eval()
already_done = load_already_done(OUTPUT_FILE)
start_idx = len(already_done)
print(f"Already translated: {start_idx}")
print(f"Output file: {OUTPUT_FILE}")
with open(OUTPUT_FILE, "a", encoding="utf-8") as out_f:
for start in range(start_idx, len(dataset), BATCH_SIZE):
end = min(start + BATCH_SIZE, len(dataset))
batch = dataset[start:end]
instructions_en = batch["instruction"]
inputs_en = batch["input"]
outputs_en = batch["output"]
instructions_sk = translate_batch(instructions_en, tokenizer, model, device)
inputs_sk = translate_batch(inputs_en, tokenizer, model, device)
outputs_sk = translate_batch(outputs_en, tokenizer, model, device)
for i in range(end - start):
row = {
"id": start + i,
"instruction_en": instructions_en[i],
"input_en": inputs_en[i],
"output_en": outputs_en[i],
"instruction_sk": instructions_sk[i],
"input_sk": inputs_sk[i],
"output_sk": outputs_sk[i],
}
out_f.write(json.dumps(row, ensure_ascii=False) + "\n")
out_f.flush()
print(f"Translated {end}/{len(dataset)}")
print("Done.")
if __name__ == "__main__":
main()