Upload folder using huggingface_hub
Browse files- app.py +43 -22
- requirements.txt +2 -0
app.py
CHANGED
|
@@ -7,12 +7,14 @@ from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
|
|
| 7 |
from fastapi.responses import JSONResponse
|
| 8 |
from huggingface_hub import HfApi
|
| 9 |
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
|
|
|
| 10 |
from pydantic import BaseModel
|
| 11 |
|
| 12 |
# ---------- Config ----------
|
| 13 |
-
|
|
|
|
| 14 |
ORG_NAME = "wolethereader"
|
| 15 |
-
|
| 16 |
LANG_CODES = {"yo": "yor_Latn", "ha": "hau_Latn", "ig": "ibo_Latn", "pcm": "pcm_Latn"}
|
| 17 |
VALID_LANGS = set(LANG_CODES.keys())
|
| 18 |
MAX_TEXT_CHARS = 2000
|
|
@@ -69,27 +71,35 @@ def check_rate_limit(username: str):
|
|
| 69 |
|
| 70 |
# ---------- Model ----------
|
| 71 |
tokenizer = None
|
| 72 |
-
model = None
|
| 73 |
|
| 74 |
@app.on_event("startup")
|
| 75 |
async def startup():
|
| 76 |
global tokenizer, model
|
| 77 |
-
log_event("
|
| 78 |
-
tokenizer = AutoTokenizer.from_pretrained(
|
| 79 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 80 |
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 81 |
-
|
|
|
|
|
|
|
|
|
|
| 82 |
model.eval()
|
| 83 |
log_event("model_loaded_ok", device=device)
|
| 84 |
|
| 85 |
class TranslateRequest(BaseModel):
|
| 86 |
text: str
|
| 87 |
-
|
|
|
|
| 88 |
max_new_tokens: int = 128
|
| 89 |
|
| 90 |
@app.get("/")
|
| 91 |
def root():
|
| 92 |
-
return {"status": "ok", "languages": sorted(VALID_LANGS), "engine":
|
| 93 |
|
| 94 |
@app.get("/health")
|
| 95 |
def health():
|
|
@@ -99,8 +109,10 @@ def health():
|
|
| 99 |
async def translate(req: TranslateRequest, username: str = Depends(verify_org_token)):
|
| 100 |
check_rate_limit(username)
|
| 101 |
|
| 102 |
-
if req.
|
| 103 |
-
raise HTTPException(status_code=400, detail=f"
|
|
|
|
|
|
|
| 104 |
if not req.text or not req.text.strip():
|
| 105 |
raise HTTPException(status_code=400, detail="text must not be empty")
|
| 106 |
if len(req.text) > MAX_TEXT_CHARS:
|
|
@@ -109,26 +121,35 @@ async def translate(req: TranslateRequest, username: str = Depends(verify_org_to
|
|
| 109 |
request_id = str(uuid.uuid4())
|
| 110 |
start = time.time()
|
| 111 |
|
| 112 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 113 |
inputs = tokenizer(req.text, return_tensors="pt", truncation=True, max_length=128).to(model.device)
|
| 114 |
-
tgt_id = tokenizer.convert_tokens_to_ids(
|
|
|
|
| 115 |
with torch.no_grad():
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
max_length=None
|
| 121 |
-
|
| 122 |
translated = tokenizer.decode(out[0], skip_special_tokens=True)
|
| 123 |
elapsed_s = round(time.time() - start, 2)
|
| 124 |
|
| 125 |
log_event("translate_ok", request_id=request_id, user=username,
|
| 126 |
-
|
| 127 |
|
| 128 |
return {
|
| 129 |
"request_id": request_id,
|
| 130 |
-
"
|
| 131 |
-
"
|
| 132 |
"translated_text": translated,
|
| 133 |
"elapsed_s": elapsed_s,
|
| 134 |
}
|
|
|
|
| 7 |
from fastapi.responses import JSONResponse
|
| 8 |
from huggingface_hub import HfApi
|
| 9 |
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
| 10 |
+
from peft import PeftModel
|
| 11 |
from pydantic import BaseModel
|
| 12 |
|
| 13 |
# ---------- Config ----------
|
| 14 |
+
BASE_MODEL_ID = "wolethereader/STORM-OS-MT-3B"
|
| 15 |
+
REVERSE_ADAPTER_ID = "wolethereader/STORM-OS-MT-3B-REVERSE"
|
| 16 |
ORG_NAME = "wolethereader"
|
| 17 |
+
EN = "eng_Latn"
|
| 18 |
LANG_CODES = {"yo": "yor_Latn", "ha": "hau_Latn", "ig": "ibo_Latn", "pcm": "pcm_Latn"}
|
| 19 |
VALID_LANGS = set(LANG_CODES.keys())
|
| 20 |
MAX_TEXT_CHARS = 2000
|
|
|
|
| 71 |
|
| 72 |
# ---------- Model ----------
|
| 73 |
tokenizer = None
|
| 74 |
+
model = None # PeftModel with the reverse adapter loaded but not necessarily active
|
| 75 |
|
| 76 |
@app.on_event("startup")
|
| 77 |
async def startup():
|
| 78 |
global tokenizer, model
|
| 79 |
+
log_event("loading_tokenizer_with_pcm_token", repo=REVERSE_ADAPTER_ID)
|
| 80 |
+
tokenizer = AutoTokenizer.from_pretrained(REVERSE_ADAPTER_ID) # has pcm_Latn added
|
| 81 |
+
|
| 82 |
+
log_event("loading_base_model", repo=BASE_MODEL_ID)
|
| 83 |
+
base_model = AutoModelForSeq2SeqLM.from_pretrained(BASE_MODEL_ID, torch_dtype=torch.bfloat16)
|
| 84 |
+
base_model.resize_token_embeddings(len(tokenizer))
|
| 85 |
+
|
| 86 |
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 87 |
+
base_model.to(device)
|
| 88 |
+
|
| 89 |
+
log_event("attaching_reverse_adapter", repo=REVERSE_ADAPTER_ID)
|
| 90 |
+
model = PeftModel.from_pretrained(base_model, REVERSE_ADAPTER_ID, adapter_name="reverse")
|
| 91 |
model.eval()
|
| 92 |
log_event("model_loaded_ok", device=device)
|
| 93 |
|
| 94 |
class TranslateRequest(BaseModel):
|
| 95 |
text: str
|
| 96 |
+
direction: str # "forward" (local -> English) or "reverse" (English -> local)
|
| 97 |
+
lang: str # the local language code, yo/ha/ig/pcm, regardless of direction
|
| 98 |
max_new_tokens: int = 128
|
| 99 |
|
| 100 |
@app.get("/")
|
| 101 |
def root():
|
| 102 |
+
return {"status": "ok", "languages": sorted(VALID_LANGS), "directions": ["forward", "reverse"], "engine": BASE_MODEL_ID}
|
| 103 |
|
| 104 |
@app.get("/health")
|
| 105 |
def health():
|
|
|
|
| 109 |
async def translate(req: TranslateRequest, username: str = Depends(verify_org_token)):
|
| 110 |
check_rate_limit(username)
|
| 111 |
|
| 112 |
+
if req.lang not in VALID_LANGS:
|
| 113 |
+
raise HTTPException(status_code=400, detail=f"lang must be one of {sorted(VALID_LANGS)}")
|
| 114 |
+
if req.direction not in ("forward", "reverse"):
|
| 115 |
+
raise HTTPException(status_code=400, detail="direction must be 'forward' or 'reverse'")
|
| 116 |
if not req.text or not req.text.strip():
|
| 117 |
raise HTTPException(status_code=400, detail="text must not be empty")
|
| 118 |
if len(req.text) > MAX_TEXT_CHARS:
|
|
|
|
| 121 |
request_id = str(uuid.uuid4())
|
| 122 |
start = time.time()
|
| 123 |
|
| 124 |
+
if req.direction == "forward":
|
| 125 |
+
src_lang, tgt_lang = LANG_CODES[req.lang], EN
|
| 126 |
+
context = model.disable_adapter()
|
| 127 |
+
else:
|
| 128 |
+
src_lang, tgt_lang = EN, LANG_CODES[req.lang]
|
| 129 |
+
model.set_adapter("reverse")
|
| 130 |
+
context = None
|
| 131 |
+
|
| 132 |
+
tokenizer.src_lang = src_lang
|
| 133 |
inputs = tokenizer(req.text, return_tensors="pt", truncation=True, max_length=128).to(model.device)
|
| 134 |
+
tgt_id = tokenizer.convert_tokens_to_ids(tgt_lang)
|
| 135 |
+
|
| 136 |
with torch.no_grad():
|
| 137 |
+
if context is not None:
|
| 138 |
+
with context:
|
| 139 |
+
out = model.generate(**inputs, forced_bos_token_id=tgt_id, max_new_tokens=req.max_new_tokens, max_length=None)
|
| 140 |
+
else:
|
| 141 |
+
out = model.generate(**inputs, forced_bos_token_id=tgt_id, max_new_tokens=req.max_new_tokens, max_length=None)
|
| 142 |
+
|
| 143 |
translated = tokenizer.decode(out[0], skip_special_tokens=True)
|
| 144 |
elapsed_s = round(time.time() - start, 2)
|
| 145 |
|
| 146 |
log_event("translate_ok", request_id=request_id, user=username,
|
| 147 |
+
direction=req.direction, lang=req.lang, elapsed_s=elapsed_s)
|
| 148 |
|
| 149 |
return {
|
| 150 |
"request_id": request_id,
|
| 151 |
+
"direction": req.direction,
|
| 152 |
+
"lang": req.lang,
|
| 153 |
"translated_text": translated,
|
| 154 |
"elapsed_s": elapsed_s,
|
| 155 |
}
|
requirements.txt
CHANGED
|
@@ -1,7 +1,9 @@
|
|
| 1 |
fastapi
|
| 2 |
uvicorn
|
| 3 |
transformers==5.15.0
|
|
|
|
| 4 |
torch
|
| 5 |
sentencepiece
|
| 6 |
accelerate
|
| 7 |
python-multipart
|
|
|
|
|
|
| 1 |
fastapi
|
| 2 |
uvicorn
|
| 3 |
transformers==5.15.0
|
| 4 |
+
peft==0.20.0
|
| 5 |
torch
|
| 6 |
sentencepiece
|
| 7 |
accelerate
|
| 8 |
python-multipart
|
| 9 |
+
# add reverse toggle, force rebuild 1787637858
|