HowWebWorks commited on
Commit
5777436
·
1 Parent(s): c9feb62

trust_remote_code=True, on tokenizer

Browse files
Files changed (1) hide show
  1. handler.py +12 -13
handler.py CHANGED
@@ -1,23 +1,22 @@
1
  from typing import Dict
2
  import torch
3
- from transformers import AutoModelForCausalLM, AutoTokenizer
4
 
5
  class EndpointHandler:
6
- """
7
- Minimal custom handler for InternLM2 / NuExtract-2-8B
8
- """
9
 
10
- def __init__(self, path: str = "./model"):
11
- # allow execution of custom model code
12
  self.tokenizer = AutoTokenizer.from_pretrained(
13
- path, trust_remote_code=True
 
14
  )
15
  self.model = AutoModelForCausalLM.from_pretrained(
16
  path,
17
- trust_remote_code=True, # ← key line
18
- torch_dtype=torch.float16, # load in fp16 to fit on one A10/T4
19
- device_map="auto" # send to GPU if available
20
- ).eval() # put in inference mode
21
 
22
  def __call__(self, data: Dict[str, str]) -> Dict[str, str]:
23
  prompt = data.get("inputs", "")
@@ -25,6 +24,6 @@ class EndpointHandler:
25
  return {"error": "No input provided."}
26
 
27
  inputs = self.tokenizer(prompt, return_tensors="pt").to(self.model.device)
28
- outputs = self.model.generate(**inputs, max_new_tokens=128)
29
- answer = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
30
  return {"generated_text": answer}
 
1
  from typing import Dict
2
  import torch
3
+ from transformers import AutoTokenizer, AutoModelForCausalLM
4
 
5
  class EndpointHandler:
6
+ """Custom handler for NuExtract-2-8B (InternLM2 based)."""
 
 
7
 
8
+ def __init__(self, path: str = "") -> None:
9
+ # ↓↓↓ allow the repo’s custom configuration & modelling code
10
  self.tokenizer = AutoTokenizer.from_pretrained(
11
+ path,
12
+ trust_remote_code=True # ← mandatory
13
  )
14
  self.model = AutoModelForCausalLM.from_pretrained(
15
  path,
16
+ trust_remote_code=True, # ← mandatory
17
+ torch_dtype=torch.float16, # fits on a 16 GB GPU
18
+ device_map="auto" # put tensors on the GPU
19
+ ).eval()
20
 
21
  def __call__(self, data: Dict[str, str]) -> Dict[str, str]:
22
  prompt = data.get("inputs", "")
 
24
  return {"error": "No input provided."}
25
 
26
  inputs = self.tokenizer(prompt, return_tensors="pt").to(self.model.device)
27
+ output_ids = self.model.generate(**inputs, max_new_tokens=128)
28
+ answer = self.tokenizer.decode(output_ids[0], skip_special_tokens=True)
29
  return {"generated_text": answer}