al3obdi's picture
fix: add @spaces.GPU + keepalive loop
933f9cc verified
Raw History Blame Contribute Delete
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}
@spaces.GPU
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)