Bayernator's picture
output_dir = /data/horst_checkpoints (direct bucket writes)
4713e2b
Raw
History Blame Contribute Delete
15.3 kB
import gradio as gr
import threading
import logging
import os
import difflib
import re
import subprocess
from datasets import load_dataset
from transformers import (
T5Tokenizer,
T5ForConditionalGeneration,
DataCollatorForSeq2Seq,
Trainer,
TrainingArguments,
TrainerCallback,
set_seed,
)
from huggingface_hub import HfApi
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
TRAINED_MODEL = "Bayernator/horst-spelling-model"
BASE_MODEL = "oliverguhr/spelling-correction-german-base"
DATASET_NAME = "Bayernator/german-spelling-data"
LOCAL_MODEL_DIR = "/data/horst_checkpoints"
BUCKET_NAME = "Bayernator/HORST-Rechtschreibbot-storage"
set_seed(42)
logger.info("=== SUCHE CHECKPOINTS IN /data ===")
if os.path.exists("/data"):
for root, dirs, files in os.walk("/data"):
logger.info(f"Pfad: {root} | Ordner: {dirs} | Dateien: {files[:5]}")
else:
logger.info("/data existiert nicht (nur im HF Space verfügbar)")
logger.info("=== STARTE MANUELLEN BUCKET-SYNC ===")
try:
subprocess.run(
["hf", "sync", "/data", f"hf://buckets/{BUCKET_NAME}"],
capture_output=True, text=True, check=True, timeout=300,
)
logger.info("Bucket-Sync von /data abgeschlossen")
except Exception as e:
logger.info(f"Bucket-Sync fehlgeschlagen (normal ausserhalb HF Space): {e}")
tokenizer = None
model = None
training_status = "Initialisiere..."
training_active = False
def try_load_model(source):
try:
tok = T5Tokenizer.from_pretrained(source)
mod = T5ForConditionalGeneration.from_pretrained(source)
return tok, mod
except Exception:
return None, None
def model_exists_on_hub(repo_id):
try:
api = HfApi()
api.model_info(repo_id, token=os.environ.get("HF_TOKEN"))
return True
except Exception:
return False
def sync_to_bucket(local_path, bucket_path):
try:
result = subprocess.run(
["hf", "sync", local_path, f"hf://buckets/{BUCKET_NAME}/{bucket_path}"],
capture_output=True, text=True, timeout=300,
)
if result.returncode == 0:
logger.info(f"Synced {local_path} to bucket {bucket_path}")
return True
else:
logger.warning(f"Bucket sync failed: {result.stderr[:200]}")
return False
except Exception as e:
logger.warning(f"Bucket sync error: {e}")
return False
class BucketSyncCallback(TrainerCallback):
def on_save(self, args, state, control, **kwargs):
step = state.global_step
logger.info(f"Checkpoint saved at step {step}. Syncing to bucket...")
sync_to_bucket(args.output_dir, f"checkpoints/step_{step}")
def word_diff(original, corrected):
orig_words = re.findall(r'\S+|\s+', original)
corr_words = re.findall(r'\S+|\s+', corrected)
orig_tokens = [w for w in orig_words if w.strip()]
corr_tokens = [w for w in corr_words if w.strip()]
matcher = difflib.SequenceMatcher(None, orig_tokens, corr_tokens)
changes = []
for tag, i1, i2, j1, j2 in matcher.get_opcodes():
if tag == "replace":
for idx in range(max(i2 - i1, j2 - j1)):
o = orig_tokens[i1 + idx] if i1 + idx < i2 else ""
c = corr_tokens[j1 + idx] if j1 + idx < j2 else ""
if o and c and o != c:
changes.append((o, c))
elif tag == "delete":
for idx in range(i1, i2):
changes.append((orig_tokens[idx], "\u2205"))
elif tag == "insert":
for idx in range(j1, j2):
changes.append(("\u2205", corr_tokens[idx]))
return changes
def build_changes_html(changes):
if not changes:
return '<div style="color: var(--text-muted); padding: 12px;">\u2714 Keine \u00c4nderungen erforderlich.</div>'
items = []
for old, new in changes:
items.append(
f'<div style="display: flex; align-items: center; gap: 10px; padding: 6px 12px; '
f'border-radius: 10px; background: rgba(139,92,246,0.08); margin-bottom: 4px;">'
f'<span style="color: #f87171; text-decoration: line-through; font-weight: 500;">{old}</span>'
f'<span style="color: var(--text-muted);">\u2192</span>'
f'<span style="color: #34d399; font-weight: 600;">{new}</span>'
f'</div>'
)
return (
f'<div style="margin-top: 12px;">'
f'<div style="font-size: 13px; color: var(--text-muted); margin-bottom: 8px; font-weight: 600;">'
f'\u00c4nderungen ({len(changes)})</div>'
f'{"".join(items)}'
f'</div>'
)
def do_training():
global model, tokenizer, training_status, training_active
if training_active:
return
training_active = True
try:
training_status = "Lade Trainingsdaten..."
logger.info("Loading dataset...")
dataset = load_dataset(DATASET_NAME, split="train")
dataset = dataset.train_test_split(test_size=0.1, seed=42)
train_data = dataset["train"]
eval_data = dataset["test"]
logger.info(f"Dataset loaded: {len(train_data)} train, {len(eval_data)} eval")
training_status = f"Tokenisiere {len(train_data)} Beispiele..."
def tokenize_fn(batch):
inputs = tokenizer(
batch["input"],
max_length=128,
truncation=True,
padding=False,
)
targets = tokenizer(
batch["target"],
max_length=128,
truncation=True,
padding=False,
)
inputs["labels"] = targets["input_ids"]
return inputs
train_data = train_data.map(
tokenize_fn, batched=True, batch_size=50,
remove_columns=["input", "target"],
)
eval_data = eval_data.map(
tokenize_fn, batched=True, batch_size=50,
remove_columns=["input", "target"],
)
train_data.set_format("torch", columns=["input_ids", "attention_mask", "labels"])
eval_data.set_format("torch", columns=["input_ids", "attention_mask", "labels"])
data_collator = DataCollatorForSeq2Seq(tokenizer, model=model, padding=True)
training_args = TrainingArguments(
output_dir=LOCAL_MODEL_DIR,
num_train_epochs=3,
per_device_train_batch_size=1,
per_device_eval_batch_size=1,
gradient_accumulation_steps=2,
learning_rate=3e-5,
warmup_steps=100,
logging_steps=5,
eval_strategy="steps",
eval_steps=500,
save_strategy="steps",
save_steps=500,
save_total_limit=2,
load_best_model_at_end=True,
metric_for_best_model="loss",
greater_is_better=False,
bf16=False,
fp16=False,
report_to="none",
dataloader_num_workers=0,
)
trainer = Trainer(
model=model,
args=training_args,
data_collator=data_collator,
train_dataset=train_data,
eval_dataset=eval_data,
processing_class=tokenizer,
callbacks=[BucketSyncCallback],
)
training_status = "Training läuft (ca. 3-5h auf CPU)..."
logger.info("Starting training...")
trainer.train()
training_status = "Speichere Modell..."
logger.info("Saving model...")
trainer.save_model(LOCAL_MODEL_DIR)
try:
training_status = "Pushe Modell zu Hugging Face..."
logger.info("Pushing to hub...")
trainer.push_to_hub(
commit_message="Trained via HORST Space on CPU",
token=os.environ.get("HF_TOKEN"),
)
except Exception as hub_error:
logger.warning(f"Push to Hub failed: {hub_error}")
try:
training_status = "Sync Modell zu Bucket..."
logger.info("Syncing to bucket...")
sync_to_bucket(LOCAL_MODEL_DIR, "model_checkpoints")
except Exception as bucket_error:
logger.warning(f"Bucket sync failed: {bucket_error}")
tokenizer = T5Tokenizer.from_pretrained(LOCAL_MODEL_DIR)
model = T5ForConditionalGeneration.from_pretrained(LOCAL_MODEL_DIR)
training_status = "Training abgeschlossen. Modell bereit."
logger.info("Training completed.")
except Exception as e:
training_status = f"Training fehlgeschlagen: {e}"
logger.error(f"Training failed: {e}", exc_info=True)
finally:
training_active = False
if os.path.isdir(LOCAL_MODEL_DIR):
tokenizer, model = try_load_model(LOCAL_MODEL_DIR)
if tokenizer is not None:
training_status = "Feinabgestimmtes Modell bereit (lokal)."
if tokenizer is None and model_exists_on_hub(TRAINED_MODEL):
tokenizer, model = try_load_model(TRAINED_MODEL)
if tokenizer is not None:
training_status = "Feinabgestimmtes Modell bereit (Hub)."
if tokenizer is None:
tokenizer = T5Tokenizer.from_pretrained(BASE_MODEL)
model = T5ForConditionalGeneration.from_pretrained(BASE_MODEL)
training_status = "Basis-Modell geladen. Training startet automatisch..."
logger.info("Base model loaded. Auto-training will start.")
def auto_train():
global training_status
training_status = "Starte automatisches Training..."
do_training()
threading.Thread(target=auto_train, daemon=True).start()
def start_training():
global training_status
if training_active:
return training_status
lower = training_status.lower()
if "lade" in lower or "läuft" in lower or "starte" in lower:
return training_status
training_status = "Training wird gestartet..."
threading.Thread(target=do_training, daemon=True).start()
return training_status
def get_training_status():
return training_status
def build_prefix(rs, gram, stil):
if rs and gram and stil:
return "Verbessere Rechtschreibung, Grammatik und Stil: "
if rs and gram:
return "Korrigiere Rechtschreibung und Grammatik: "
if rs and stil:
return "Korrigiere Rechtschreibung und verbessere den Stil: "
if gram and stil:
return "Korrigiere Grammatik und verbessere den Stil: "
if gram:
return "Korrigiere die Grammatik: "
if stil:
return "Verbessere den Stil: "
return "Korrigiere die Rechtschreibung: "
def format_info(text_len, prefix):
if "Verbessere" in prefix and "Grammatik" in prefix and "Stil" in prefix:
mode = "RS + Grammatik + Stil"
elif "Grammatik" in prefix and "Stil" in prefix:
mode = "Grammatik + Stil"
elif "Rechtschreibung" in prefix and "Stil" in prefix:
mode = "RS + Stil"
elif "Grammatik" in prefix:
mode = "Grammatik"
elif "Stil" in prefix:
mode = "Stil"
else:
mode = "Rechtschreibung"
return f"{text_len} Zeichen \u00b7 {mode}"
def korrigieren(text, rs, gram, stil):
if not text or not text.strip():
return "", "", "", "Bitte gib zuerst einen Text ein."
try:
prefix = build_prefix(rs, gram, stil)
prompt = prefix + text.strip()
original = text.strip()
inputs = tokenizer(
prompt, return_tensors="pt",
max_length=512, truncation=True,
)
outputs = model.generate(
**inputs, max_length=512,
num_beams=4, early_stopping=True,
)
result = tokenizer.decode(outputs[0], skip_special_tokens=True)
info = format_info(len(text), prefix)
changes = word_diff(original, result)
changes_html = build_changes_html(changes)
return result, changes_html, "", info
except Exception as e:
return "", "", "", f"Fehler: {e}"
CSS = """
footer { display: none !important; }
header { display: none !important; }
.gradio-container {
max-width: 960px !important;
margin: auto !important;
}
.gradio-container textarea {
border-radius: 12px !important;
font-size: 15px !important;
line-height: 1.6 !important;
padding: 14px !important;
}
.panel-label {
font-size: 12px; font-weight: 600;
color: #6b7280; letter-spacing: 0.5px;
margin-bottom: 6px;
}
"""
with gr.Blocks() as demo:
gr.HTML("""
<div style="text-align: center; padding: 24px 0 4px 0;">
<span style="font-size: 22px; font-weight: 700; color: #1f2937;">HORST</span>
<div style="font-size: 13px; color: #9ca3af; margin-top: 2px;">
KI-Schreibassistent f\u00fcr die deutsche Sprache
</div>
</div>
""")
with gr.Row(equal_height=True):
with gr.Column(scale=1, min_width=300):
gr.HTML('<div class="panel-label">EINGABE</div>')
text_input = gr.Textbox(
label=None, placeholder="Deutschen Text hier eingeben...",
lines=16, max_lines=30,
container=False,
)
with gr.Column(scale=1, min_width=300):
gr.HTML('<div class="panel-label">AUSGABE</div>')
text_output = gr.Textbox(
label=None, lines=16, max_lines=30,
interactive=False, container=False,
)
with gr.Row():
rs_cb = gr.Checkbox(label="Rechtschreibung", value=True)
gram_cb = gr.Checkbox(label="Grammatik", value=False)
stil_cb = gr.Checkbox(label="Stil", value=False)
beispiel_btn = gr.Button("\u2b50 Beispiel", variant="secondary", size="sm")
leeren_btn = gr.Button("\u21ba Leeren", variant="secondary", size="sm")
korrigieren_btn = gr.Button("\u2192 Korrigieren", variant="primary", size="sm")
changes_box = gr.HTML()
def on_correction(text, rs, gram, stil):
res, changes_html, _, info = korrigieren(text, rs, gram, stil)
return res, changes_html
def on_beispiel():
return "Ich habe gestern ein interesanten Artikel gelessen und fand ihn zimlich gut."
def on_leeren():
return "", "", ""
korrigieren_btn.click(
fn=on_correction,
inputs=[text_input, rs_cb, gram_cb, stil_cb],
outputs=[text_output, changes_box],
)
beispiel_btn.click(fn=on_beispiel, inputs=None, outputs=text_input)
leeren_btn.click(fn=on_leeren, inputs=None, outputs=[text_input, text_output, changes_box])
gr.HTML("<br>")
with gr.Accordion("\u2699 Training & Status", open=False):
training_status_box = gr.Textbox(
label="Status", value=training_status,
interactive=False, lines=2,
)
with gr.Row():
refresh_status_btn = gr.Button("Aktualisieren", size="sm")
train_btn = gr.Button("Training starten", variant="secondary", size="sm")
train_btn.click(fn=start_training, inputs=None, outputs=training_status_box)
refresh_status_btn.click(fn=get_training_status, inputs=None, outputs=training_status_box)
demo.launch(
css=CSS,
theme=gr.themes.Base(
primary_hue="blue",
neutral_hue="gray",
font=gr.themes.GoogleFont("Inter"),
),
)