AutoTranslateAI / translator_utils.py
Brian045's picture
Update translator_utils.py
636950b verified
Raw History Blame Contribute Delete
8.42 kB
import os
import pandas as pd
from transformers import pipeline
import torch
from huggingface_hub import HfApi, hf_hub_download
from llama_index.core.node_parser import SentenceSplitter
# --- Configuration & Globals ---
DATA_FILE = "data_berita.csv"
ENV_HF_TOKEN = os.getenv("HF_TOKEN", "")
ENV_REPO_ID = os.getenv("DATASET_REPO_ID", "Brian045/data_berita")
DATASET_URL = f"https://huggingface.co/datasets/{ENV_REPO_ID}/resolve/main/{DATA_FILE}"
LOCAL_FILENAME = DATA_FILE
# Global model cache
_translator = None
_splitter = None
def download_dataset(force=False):
"""
Downloads the dataset if it doesn't exist locally or if forced.
Uses hf_hub_download to ensure raw file integrity (avoids parsing errors).
"""
if force and os.path.exists(LOCAL_FILENAME):
try:
os.remove(LOCAL_FILENAME)
except:
pass
if not os.path.exists(LOCAL_FILENAME):
print(f"Downloading dataset from {ENV_REPO_ID}...")
try:
# Use hf_hub_download to get the raw file
downloaded_path = hf_hub_download(
repo_id=ENV_REPO_ID,
filename=DATA_FILE,
repo_type="dataset",
token=ENV_HF_TOKEN if ENV_HF_TOKEN else None,
local_dir=".",
local_dir_use_symlinks=False
)
print(f"Download complete: {downloaded_path}")
except Exception as e:
print(f"Error downloading dataset: {e}")
pass
# Check if file is valid
if os.path.exists(LOCAL_FILENAME):
try:
if os.path.getsize(LOCAL_FILENAME) < 10:
print("File too small, removing...")
os.remove(LOCAL_FILENAME)
except:
pass
def read_csv_robust(filepath):
"""
Attempts to read CSV with various delimiters and robust settings.
"""
if not os.path.exists(filepath):
return pd.DataFrame()
separators = [',', ';', '\t', '|']
for sep in separators:
try:
df = pd.read_csv(filepath, sep=sep, on_bad_lines='skip')
if len(df.columns) > 1:
return df
except:
continue
# Fallback: Python engine
try:
df = pd.read_csv(filepath, sep=None, engine='python', on_bad_lines='skip')
return df
except Exception as e:
print(f"Failed to read CSV: {e}")
return pd.DataFrame()
def upload_dataset(token):
"""
Uploads the local CSV file back to the Hugging Face Hub.
"""
token_to_use = token if token else ENV_HF_TOKEN
if not token_to_use:
print("No HF Token provided. Skipping upload.")
return "Skipped Upload (No Token)"
print(f"Uploading {LOCAL_FILENAME} to {ENV_REPO_ID}...")
try:
api = HfApi()
api.upload_file(
path_or_fileobj=LOCAL_FILENAME,
path_in_repo=DATA_FILE,
repo_id=ENV_REPO_ID,
repo_type="dataset",
token=token_to_use,
commit_message="Update translations via Auto AI Translator"
)
print("Upload successful!")
return "Upload Successful!"
except Exception as e:
print(f"Error uploading dataset: {e}")
return f"Upload Failed: {str(e)}"
def load_resources():
"""
Loads the NLLB-200 translation pipeline and LlamaIndex splitter.
"""
global _translator, _splitter
if _translator is None:
print("Loading NLLB-200 model...")
device = 0 if torch.cuda.is_available() else -1
# NLLB requires source/target langs.
# ind_Latn = Indonesian, eng_Latn = English
_translator = pipeline("translation", model="facebook/nllb-200-1.3B", src_lang="ind_Latn", tgt_lang="eng_Latn", device=device)
if _splitter is None:
print("Loading LlamaIndex SentenceSplitter...")
# Reduce chunk_size to 64 to ensure NLLB doesn't truncate output.
# NLLB seems to struggle with long context > 200 tokens output generation.
# 64 tokens is safer to ensure complete translation of every sentence.
_splitter = SentenceSplitter(chunk_size=64, chunk_overlap=0)
return _translator, _splitter
def smart_translate(text, translator, splitter):
"""
Translates text using LlamaIndex SentenceSplitter and NLLB-200.
"""
if not text or not isinstance(text, str) or text.strip() == "":
return ""
# 1. Chunking via LlamaIndex
chunks = splitter.split_text(text)
# 2. Translation
translated_parts = []
for chunk in chunks:
if not chunk.strip(): continue
try:
# max_length=512 is plenty for a 128 token input
res = translator(chunk, max_length=512, truncation=True)
translated_parts.append(res[0]['translation_text'])
except Exception as e:
print(f"Chunk translation error: {e}")
translated_parts.append(chunk)
return " ".join(translated_parts)
def process_rows(indices, token=None, progress=None):
"""
Translates only the specified row indices.
"""
download_dataset()
df = read_csv_robust(LOCAL_FILENAME)
if df.empty:
return df, "Error: Dataset empty or failed to load."
if 'Judul_Inggris' not in df.columns:
df['Judul_Inggris'] = ""
if 'Isi_Inggris' not in df.columns:
df['Isi_Inggris'] = ""
df['Judul_Inggris'] = df['Judul_Inggris'].fillna("")
df['Isi_Inggris'] = df['Isi_Inggris'].fillna("")
df['Judul'] = df['Judul'].astype(str)
df['Isi'] = df['Isi'].astype(str)
if not indices:
return df, "No rows selected."
translator, splitter = load_resources()
total = len(indices)
print(f"Processing {total} selected rows...")
token_arg = token if token and token.strip() else ENV_HF_TOKEN
api = HfApi() if token_arg else None
count = 0
last_upload_status = "No upload (token missing)"
for idx in indices:
if idx not in df.index:
continue
row = df.loc[idx]
# Translate Judul
if not row['Judul_Inggris'] or row['Judul_Inggris'].strip() == "":
df.at[idx, 'Judul_Inggris'] = smart_translate(row['Judul'], translator, splitter)
# Translate Isi
if not row['Isi_Inggris'] or row['Isi_Inggris'].strip() == "":
df.at[idx, 'Isi_Inggris'] = smart_translate(row['Isi'], translator, splitter)
# INCREMENTAL SAVE & UPLOAD
# 1. Save locally immediately
df.to_csv(LOCAL_FILENAME, index=False)
# 2. Upload immediately (if token exists)
if api:
try:
if progress:
progress((count + 0.5) / total, desc=f"Uploading row {idx}...")
api.upload_file(
path_or_fileobj=LOCAL_FILENAME,
path_in_repo=DATA_FILE,
repo_id=ENV_REPO_ID,
repo_type="dataset",
token=token_arg,
commit_message=f"Auto AI Translate: Row {idx}"
)
last_upload_status = "Upload Successful!"
except Exception as e:
print(f"Intermediate upload failed for row {idx}: {e}")
last_upload_status = f"Partial Upload Fail: {str(e)}"
count += 1
if progress:
progress(count / total, desc=f"Completed row {idx}")
return df, f"Processed {count} rows. Last status: {last_upload_status}"
def revert_rows(indices, token=None, progress=None):
"""
Reverts translations for the specified row indices.
"""
download_dataset()
df = read_csv_robust(LOCAL_FILENAME)
if not indices:
return df, "No rows selected."
print(f"Reverting {len(indices)} rows...")
for idx in indices:
if idx in df.index:
df.at[idx, 'Judul_Inggris'] = ""
df.at[idx, 'Isi_Inggris'] = ""
print("Saving to CSV...")
df.to_csv(LOCAL_FILENAME, index=False)
token_arg = token if token and token.strip() else None
upload_status = upload_dataset(token_arg)
return df, f"Reverted {len(indices)} rows. {upload_status}"