| import os |
| import time |
| import json |
| import gradio as gr |
|
|
| def get_token(): |
| return os.environ.get("HF_TOKEN", os.environ.get("HUGGING_FACE_HUB_TOKEN", "")) |
|
|
| def train_model(): |
| HF_TOKEN = get_token() |
| if not HF_TOKEN: |
| yield "HATA: HF_TOKEN yok!", "HATA", "", "" |
| return |
|
|
| logs = [] |
| def log(msg): |
| t = time.strftime("%H:%M:%S") |
| logs.append(f"[{t}] {msg}") |
| return "\n".join(logs[-40:]) |
|
|
| try: |
| yield log("1/5 Dataset yukleniyor..."), "Baslatiliyor...", "", "" |
|
|
| from datasets import load_dataset |
| ds = load_dataset("remox01/xde", split="train", token=HF_TOKEN) |
| yield log(f" {len(ds)} cumle"), "Hazir", "", "" |
|
|
| yield log("2/5 Tokenizer..."), "Hazir", "", "" |
| from transformers import AutoTokenizer, AutoModelForSeq2SeqLM |
| tokenizer = AutoTokenizer.from_pretrained("Helsinki-NLP/opus-mt-tc-big-tr-en") |
| tokenizer.pad_token = tokenizer.eos_token |
|
|
| yield log("3/5 Model..."), "Hazir", "", "" |
| model = AutoModelForSeq2SeqLM.from_pretrained("Helsinki-NLP/opus-mt-tc-big-tr-en") |
| from peft import LoraConfig, get_peft_model |
| model = get_peft_model(model, LoraConfig(r=8, lora_alpha=16, target_modules=["q_proj","v_proj"], lora_dropout=0.1, bias="none", task_type="SEQ_2_SEQ_LM")) |
|
|
| yield log("4/5 Veri..."), "Hazir", "", "" |
| def preprocess(ex): |
| inp = tokenizer(ex["turkish"], max_length=128, truncation=True, padding="max_length") |
| lbl = tokenizer(ex["english"], max_length=128, truncation=True, padding="max_length") |
| lbl["input_ids"] = [[-100 if t==tokenizer.pad_token_id else t for t in l] for l in lbl["input_ids"]] |
| return {"input_ids": inp["input_ids"], "attention_mask": inp["attention_mask"], "labels": lbl["input_ids"]} |
| tokenized = ds.map(preprocess, batched=True) |
|
|
| yield log(f"5/5 Egitim basliyor (CPU - yavas olacak)..."), "Egitimde...", "", "CPU'da egitim yavas olacak. T4 GPU secin!" |
|
|
| from transformers import Seq2SeqTrainingArguments, Seq2SeqTrainer, DataCollatorForSeq2Seq |
| args = Seq2SeqTrainingArguments(output_dir="./r", per_device_train_batch_size=16, num_train_epochs=1, |
| learning_rate=2e-4, logging_steps=5, save_strategy="no", push_to_hub=False, report_to="none") |
|
|
| all_loss = [] |
| class C: |
| def on_init_end(self, *a, **k): pass |
| def on_train_begin(self, *a, **k): pass |
| def on_train_end(self, *a, **k): pass |
| def on_epoch_begin(self, *a, **k): pass |
| def on_epoch_end(self, *a, **k): pass |
| def on_step_begin(self, *a, **k): pass |
| def on_step_end(self, *a, **k): pass |
| def on_save(self, *a, **k): pass |
| def on_evaluate(self, *a, **k): pass |
| def on_prediction_step(self, *a, **k): pass |
| def on_log(self, args, state, control, logs=None, **kw): |
| if logs and "loss" in logs: |
| all_loss.append({"step": state.global_step, "loss": round(logs["loss"],4), "epoch": round(logs.get("epoch",0),2)}) |
|
|
| trainer = Seq2SeqTrainer(model=model, args=args, train_dataset=tokenized, |
| data_collator=DataCollatorForSeq2Seq(tokenizer, model=model), callbacks=[C()]) |
|
|
| start = time.time() |
| trainer.train() |
| elapsed = time.time() - start |
|
|
| losses = [x["loss"] for x in all_loss] |
| final = losses[-1] if losses else 0 |
| mini = min(losses) if losses else 0 |
| first = losses[0] if losses else 0 |
| pct = ((first-final)/first*100) if first else 0 |
|
|
| summary = f"Sure: {elapsed:.0f}s ({elapsed/60:.1f}dk)\nStep: {len(losses)}\nSon: {final}\nMin: {mini}\nAzalim: %{pct:.1f}" |
| mjson = json.dumps(all_loss, indent=2) |
|
|
| yield log(f" BITTI! {elapsed:.0f}s"), f"BITTI | {elapsed:.0f}s", mjson, summary |
|
|
| except Exception as e: |
| yield log(f" HATA: {e}"), "HATA", "", str(e) |
|
|
| with gr.Blocks(title="XDE Egitim") as demo: |
| gr.Markdown("# XDE Egitim (CPU - yavas)\nT4 GPU sec: Settings > Hardware > T4 small") |
| with gr.Row(): |
| btn = gr.Button("EGITIMI BASLAT", variant="primary") |
| status = gr.Textbox(label="Durum") |
| with gr.Row(): |
| log_box = gr.Textbox(label="Log", lines=20) |
| with gr.Column(): |
| m_box = gr.Textbox(label="Metrikler", lines=10) |
| s_box = gr.Textbox(label="Ozet", lines=10) |
| btn.click(train_model, outputs=[log_box, status, m_box, s_box]) |
| demo.launch(server_name="0.0.0.0", server_port=7860) |
|
|