Spaces:
Sleeping
Sleeping
| 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"), | |
| ), | |
| ) | |