ememzyvisuals commited on
Commit
27a33aa
·
verified ·
1 Parent(s): 89cbf60

Upload folder using huggingface_hub

Browse files
Files changed (2) hide show
  1. app.py +43 -22
  2. 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
- MODEL_ID = "wolethereader/STORM-OS-MT-3B"
 
14
  ORG_NAME = "wolethereader"
15
- TGT_LANG = "eng_Latn"
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("loading_model", model=MODEL_ID)
78
- tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
79
- model = AutoModelForSeq2SeqLM.from_pretrained(MODEL_ID, torch_dtype=torch.bfloat16)
 
 
 
 
80
  device = "cuda" if torch.cuda.is_available() else "cpu"
81
- model.to(device)
 
 
 
82
  model.eval()
83
  log_event("model_loaded_ok", device=device)
84
 
85
  class TranslateRequest(BaseModel):
86
  text: str
87
- source_lang: str
 
88
  max_new_tokens: int = 128
89
 
90
  @app.get("/")
91
  def root():
92
- return {"status": "ok", "languages": sorted(VALID_LANGS), "engine": MODEL_ID}
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.source_lang not in VALID_LANGS:
103
- raise HTTPException(status_code=400, detail=f"source_lang must be one of {sorted(VALID_LANGS)}")
 
 
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
- tokenizer.src_lang = LANG_CODES[req.source_lang]
 
 
 
 
 
 
 
 
113
  inputs = tokenizer(req.text, return_tensors="pt", truncation=True, max_length=128).to(model.device)
114
- tgt_id = tokenizer.convert_tokens_to_ids(TGT_LANG)
 
115
  with torch.no_grad():
116
- out = model.generate(
117
- **inputs,
118
- forced_bos_token_id=tgt_id,
119
- max_new_tokens=req.max_new_tokens,
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
- source_lang=req.source_lang, elapsed_s=elapsed_s)
127
 
128
  return {
129
  "request_id": request_id,
130
- "source_lang": req.source_lang,
131
- "target_lang": "en",
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