Download app.py from wolethereader/STORM-OS-MT-API: direct link, hf CLI and curl.
- Browser
- Download file 6.33 kB
-
https://huggingface.co/spaces/wolethereader/STORM-OS-MT-API/resolve/main/app.py
- Command line
-
hf download hf://spaces/wolethereader/STORM-OS-MT-API/app.py
-
curl -L -o app.py https://huggingface.co/spaces/wolethereader/STORM-OS-MT-API/resolve/main/app.py
6.33 kB
| import io, json, logging, os, time, uuid, asyncio | |
| from collections import defaultdict, deque | |
| import torch | |
| from fastapi import FastAPI, HTTPException, Depends, Request | |
| from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials | |
| from fastapi.responses import JSONResponse | |
| from huggingface_hub import HfApi | |
| from transformers import AutoModelForSeq2SeqLM, AutoTokenizer | |
| from peft import PeftModel | |
| from pydantic import BaseModel | |
| # ---------- Config ---------- | |
| BASE_MODEL_ID = "wolethereader/STORM-OS-MT-3B" | |
| REVERSE_ADAPTER_ID = "wolethereader/STORM-OS-MT-3B-REVERSE" | |
| HF_TOKEN = os.environ.get("HF_TOKEN") # required: both model repos are private | |
| ORG_NAME = "wolethereader" | |
| EN = "eng_Latn" | |
| LANG_CODES = {"yo": "yor_Latn", "ha": "hau_Latn", "ig": "ibo_Latn", "pcm": "pcm_Latn"} | |
| VALID_LANGS = set(LANG_CODES.keys()) | |
| MAX_TEXT_CHARS = 2000 | |
| app = FastAPI(title="STORM-OS MT API") | |
| def log_event(event, **fields): | |
| print(json.dumps({"event": event, "ts": time.time(), **fields})) | |
| # ---------- Auth: HF token AND must belong to the org ---------- | |
| security = HTTPBearer() | |
| hf_api = HfApi() | |
| _token_cache = {} | |
| TOKEN_CACHE_TTL = 300 | |
| EXTERNAL_ACCESS_TOKEN = os.environ.get("EXTERNAL_ACCESS_TOKEN") | |
| async def verify_org_token(creds: HTTPAuthorizationCredentials = Depends(security)): | |
| token = creds.credentials | |
| if EXTERNAL_ACCESS_TOKEN and token == EXTERNAL_ACCESS_TOKEN: | |
| log_event("auth_external_token_used") | |
| return "external-collaborator" | |
| now = time.time() | |
| cached = _token_cache.get(token) | |
| if cached and cached[1] > now: | |
| return cached[0] | |
| try: | |
| info = hf_api.whoami(token=token) | |
| except Exception: | |
| log_event("auth_failed_invalid_token") | |
| raise HTTPException(status_code=401, detail="Invalid or expired Hugging Face token") | |
| username = info.get("name", "unknown") | |
| user_orgs = [o.get("name") for o in info.get("orgs", [])] | |
| if ORG_NAME not in user_orgs: | |
| log_event("auth_failed_not_org_member", user=username, orgs=user_orgs) | |
| raise HTTPException(status_code=403, detail=f"Token does not belong to a member of '{ORG_NAME}'") | |
| _token_cache[token] = (username, now + TOKEN_CACHE_TTL) | |
| return username | |
| # ---------- Rate limiting ---------- | |
| _rate_state = defaultdict(deque) | |
| RATE_LIMIT_PER_MIN = 30 | |
| def check_rate_limit(username: str): | |
| now = time.time() | |
| q = _rate_state[username] | |
| while q and q[0] < now - 60: | |
| q.popleft() | |
| if len(q) >= RATE_LIMIT_PER_MIN: | |
| raise HTTPException(status_code=429, detail="Rate limit exceeded, try again shortly") | |
| q.append(now) | |
| # ---------- Model ---------- | |
| tokenizer = None | |
| model = None # PeftModel with the reverse adapter loaded but not necessarily active | |
| async def startup(): | |
| global tokenizer, model | |
| log_event("loading_tokenizer_with_pcm_token", repo=REVERSE_ADAPTER_ID) | |
| tokenizer = AutoTokenizer.from_pretrained(REVERSE_ADAPTER_ID, token=HF_TOKEN) # has pcm_Latn added | |
| log_event("loading_base_model", repo=BASE_MODEL_ID) | |
| base_model = AutoModelForSeq2SeqLM.from_pretrained(BASE_MODEL_ID, torch_dtype=torch.bfloat16, token=HF_TOKEN) | |
| base_model.resize_token_embeddings(len(tokenizer)) | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| base_model.to(device) | |
| log_event("attaching_reverse_adapter", repo=REVERSE_ADAPTER_ID) | |
| model = PeftModel.from_pretrained(base_model, REVERSE_ADAPTER_ID, adapter_name="reverse", token=HF_TOKEN) | |
| model.eval() | |
| log_event("model_loaded_ok", device=device) | |
| class TranslateRequest(BaseModel): | |
| text: str | |
| direction: str # "forward" (local -> English) or "reverse" (English -> local) | |
| lang: str # the local language code, yo/ha/ig/pcm, regardless of direction | |
| max_new_tokens: int = 128 | |
| def root(): | |
| return {"status": "ok", "languages": sorted(VALID_LANGS), "directions": ["forward", "reverse"], "engine": BASE_MODEL_ID} | |
| def health(): | |
| return {"status": "ok" if model is not None else "loading"} | |
| async def translate(req: TranslateRequest, username: str = Depends(verify_org_token)): | |
| check_rate_limit(username) | |
| if req.lang not in VALID_LANGS: | |
| raise HTTPException(status_code=400, detail=f"lang must be one of {sorted(VALID_LANGS)}") | |
| if req.direction not in ("forward", "reverse"): | |
| raise HTTPException(status_code=400, detail="direction must be 'forward' or 'reverse'") | |
| if not req.text or not req.text.strip(): | |
| raise HTTPException(status_code=400, detail="text must not be empty") | |
| if len(req.text) > MAX_TEXT_CHARS: | |
| raise HTTPException(status_code=400, detail=f"text exceeds {MAX_TEXT_CHARS} character limit") | |
| request_id = str(uuid.uuid4()) | |
| start = time.time() | |
| if req.direction == "forward": | |
| src_lang, tgt_lang = LANG_CODES[req.lang], EN | |
| context = model.disable_adapter() | |
| else: | |
| src_lang, tgt_lang = EN, LANG_CODES[req.lang] | |
| model.set_adapter("reverse") | |
| context = None | |
| tokenizer.src_lang = src_lang | |
| inputs = tokenizer(req.text, return_tensors="pt", truncation=True, max_length=128).to(model.device) | |
| tgt_id = tokenizer.convert_tokens_to_ids(tgt_lang) | |
| with torch.no_grad(): | |
| if context is not None: | |
| with context: | |
| out = model.generate(**inputs, forced_bos_token_id=tgt_id, max_new_tokens=req.max_new_tokens, max_length=None) | |
| else: | |
| out = model.generate(**inputs, forced_bos_token_id=tgt_id, max_new_tokens=req.max_new_tokens, max_length=None) | |
| translated = tokenizer.decode(out[0], skip_special_tokens=True) | |
| elapsed_s = round(time.time() - start, 2) | |
| log_event("translate_ok", request_id=request_id, user=username, | |
| direction=req.direction, lang=req.lang, elapsed_s=elapsed_s) | |
| return { | |
| "request_id": request_id, | |
| "direction": req.direction, | |
| "lang": req.lang, | |
| "translated_text": translated, | |
| "elapsed_s": elapsed_s, | |
| } | |
| async def http_exception_handler(request: Request, exc: HTTPException): | |
| log_event("request_error", path=str(request.url.path), status_code=exc.status_code, detail=exc.detail) | |
| return JSONResponse(status_code=exc.status_code, content={"detail": exc.detail}) | |