174 lines
4.4 KiB
Python
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()
|