Download app.py from al3obdi/thaqafa-repe-extraction: direct link, hf CLI and curl.
- Browser
- Download file 5.96 kB
-
https://huggingface.co/spaces/al3obdi/thaqafa-repe-extraction/resolve/main/app.py
- Command line
-
hf download hf://spaces/al3obdi/thaqafa-repe-extraction/app.py
-
curl -L -o app.py https://huggingface.co/spaces/al3obdi/thaqafa-repe-extraction/resolve/main/app.py
5.96 kB
| """Thaqafa-RepE Zero-GPU Vector Extraction Space.""" | |
| import os | |
| import json | |
| import traceback | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from typing import Any | |
| os.environ["GRADIO_SERVER_NAME"] = "0.0.0.0" | |
| os.environ["GRADIO_SERVER_PORT"] = "7860" | |
| import gradio as gr | |
| import torch | |
| import spaces # HF ZeroGPU SDK | |
| DEFAULT_DATASET = "al3obdi/thaqafa-repe-vectors" | |
| DEFAULT_CONCEPTS = "wasta_001, muruah_001, diyafa_001" | |
| MODEL_CHOICES = [ | |
| "meta-llama/Meta-Llama-3-8B-Instruct", | |
| "allam-ai/ALLaM-1-7b-Instruct", | |
| "core42/jais-13b-chat", | |
| ] | |
| def _resolve_token(token=None): | |
| t = token or os.environ.get("HF_TOKEN") | |
| if not t: | |
| raise ValueError("No HF_TOKEN set") | |
| return t | |
| def _tensor_to_list(tensor): | |
| return tensor.detach().cpu().to(torch.float32).tolist() | |
| def _list_to_tensor(values): | |
| return torch.tensor(values, dtype=torch.float32) | |
| def _save_to_hf(vectors, model_name, extraction_layers, dataset_name): | |
| from datasets import Dataset | |
| token = _resolve_token() | |
| rows = [] | |
| ts = datetime.now(timezone.utc).isoformat() | |
| for cid, vec in vectors.items(): | |
| rows.append({ | |
| "concept_id": cid, "concept_ar": "", "concept_en": "", | |
| "vector": _tensor_to_list(vec), | |
| "extraction_layer": extraction_layers.get(cid, -1), | |
| "model_name": model_name, "extraction_timestamp": ts, | |
| }) | |
| ds = Dataset.from_list(rows) | |
| ds.push_to_hub(dataset_name, token=token, private=True) | |
| return f"https://huggingface.co/datasets/{dataset_name}" | |
| def _load_from_hf(dataset_name): | |
| from datasets import load_dataset | |
| token = _resolve_token() | |
| ds = load_dataset(dataset_name, token=token) | |
| split = list(ds.keys())[0] if hasattr(ds, "keys") else "train" | |
| rows = [dict(r) for r in ds[split]] | |
| return {r["concept_id"]: _list_to_tensor(r["vector"]) for r in rows} | |
| def extract_vectors(concept_ids, model_name, dataset_name, progress=None): | |
| if progress is None: | |
| progress = gr.Progress() | |
| if not concept_ids.strip(): | |
| return "Please provide concept IDs." | |
| ids = [c.strip() for c in concept_ids.split(",") if c.strip()] | |
| if not ids: | |
| return "No valid concept IDs." | |
| progress(0.1, desc="Loading model...") | |
| try: | |
| from transformer_lens import HookedTransformer | |
| hf_token = os.environ.get("HF_TOKEN") | |
| kwargs = {"token": hf_token} if hf_token else {} | |
| model = HookedTransformer.from_pretrained( | |
| model_name, device="cuda", dtype=torch.bfloat16, **kwargs | |
| ) | |
| model.eval() | |
| except Exception as e: | |
| return f"Model load failed: {e}\n{traceback.format_exc()}" | |
| n_layers = int(model.cfg.n_layers) | |
| layer = n_layers // 2 | |
| hook = f"blocks.{layer}.hook_resid_post" | |
| progress(0.3, desc="Extracting...") | |
| results = {} | |
| layers_out = {} | |
| for i, cid in enumerate(ids): | |
| progress(0.3 + 0.5 * (i+1)/len(ids), desc=f"Extracting {cid}...") | |
| try: | |
| with torch.no_grad(): | |
| tokens = model.to_tokens([f"Concept: {cid}"]) | |
| _, cache = model.run_with_cache( | |
| tokens, names_filter=hook, | |
| stop_at_layer=layer+1, return_type=None, | |
| ) | |
| vec = cache[hook][0].mean(dim=0).to(torch.float32).cpu() | |
| vec = vec / vec.norm() | |
| results[cid] = vec | |
| layers_out[cid] = layer | |
| except Exception as e: | |
| return f"Extraction failed for {cid}: {e}" | |
| progress(0.85, desc="Saving to HF...") | |
| try: | |
| url = _save_to_hf(results, model_name, layers_out, dataset_name) | |
| except Exception as e: | |
| return f"Push failed: {e}" | |
| progress(1.0, desc="Done!") | |
| return f"Extracted {len(results)} vectors\nPushed to: {url}" | |
| def preview_results(dataset_name): | |
| try: | |
| vectors = _load_from_hf(dataset_name) | |
| except Exception as e: | |
| return json.dumps({"error": str(e)}, indent=2) | |
| preview = {} | |
| for cid, vec in list(vectors.items())[:5]: | |
| preview[cid] = {"shape": list(vec.shape), "norm": float(vec.norm().item())} | |
| return json.dumps(preview, indent=2, ensure_ascii=False) | |
| def download_json(dataset_name): | |
| vectors = _load_from_hf(dataset_name) | |
| output = {} | |
| for cid, vec in vectors.items(): | |
| output[cid] = {"vector": vec.tolist(), "shape": list(vec.shape)} | |
| p = "/tmp/thaqafa_vectors.json" | |
| Path(p).write_text(json.dumps(output, indent=2, ensure_ascii=False)) | |
| return p | |
| def push_to_hf(dataset_name): | |
| return f"Auto-pushed after extraction. Dataset: {dataset_name}" | |
| demo = gr.Blocks(title="Thaqafa-RepE") | |
| with demo: | |
| gr.Markdown("# Thaqafa-RepE Zero-GPU Vector Extraction") | |
| with gr.Tab("Extract Vectors"): | |
| concept_input = gr.Textbox(label="Concept IDs", value=DEFAULT_CONCEPTS) | |
| model_dropdown = gr.Dropdown(choices=MODEL_CHOICES, value=MODEL_CHOICES[0], label="Model") | |
| dataset_input = gr.Textbox(label="HF Dataset", value=DEFAULT_DATASET) | |
| extract_btn = gr.Button("Extract & Push", variant="primary") | |
| extract_output = gr.Textbox(label="Status", lines=6) | |
| with gr.Tab("Preview Results"): | |
| preview_btn = gr.Button("Load Latest Results") | |
| preview_output = gr.Code(label="Preview", language="json") | |
| with gr.Tab("Download"): | |
| download_btn = gr.Button("Download JSON") | |
| download_file = gr.File(label="File") | |
| push_btn = gr.Button("Push to HF") | |
| push_output = gr.Textbox(label="Status", lines=2) | |
| extract_btn.click(extract_vectors, [concept_input, model_dropdown, dataset_input], extract_output) | |
| preview_btn.click(preview_results, [dataset_input], preview_output) | |
| download_btn.click(download_json, [dataset_input], download_file) | |
| push_btn.click(push_to_hf, [dataset_input], push_output) | |
| demo.queue(default_concurrency_limit=1) | |
| demo.launch() | |
| import time | |
| while True: | |
| time.sleep(3600) | |