lxz8798 commited on
Commit
871edca
·
verified ·
1 Parent(s): 5b71e1d

Update handler.py

Browse files
Files changed (1) hide show
  1. handler.py +17 -56
handler.py CHANGED
@@ -1,68 +1,29 @@
1
  import subprocess
2
  import sys
 
3
 
4
- # 🚀 强制在代码运行最开始,从 PyPI 热升级 transformers 和 accelerate 到最新版本
5
  try:
6
- print("Starting hot-upgrade of transformers and accelerate...")
7
  subprocess.run(
8
  [sys.executable, "-m", "pip", "install", "--upgrade", "transformers", "accelerate", "--user"],
9
  check=True
10
  )
11
- print("Hot-upgrade completed successfully!")
12
  except Exception as e:
13
- print(f"Warning during hot-upgrade: {e}")
14
 
15
- # ----------------- 升级完成后再导 -----------------
16
- import torch
17
- from transformers import AutoModelForCausalLM, AutoTokenizer
18
-
19
- class EndpointHandler:
20
- def __init__(self, path=""):
21
- # 1. 使用 trust_remote_code=True 加载 Tokenizer 和 Model
22
- self.tokenizer = AutoTokenizer.from_pretrained(path, trust_remote_code=True)
23
- self.model = AutoModelForCausalLM.from_pretrained(
24
- path,
25
- torch_dtype=torch.bfloat16,
26
- device_map="auto",
27
- trust_remote_code=True
28
- )
29
 
30
- def __call__(self, data):
31
- # 2. 获取输入和参数
32
- inputs = data.get("inputs", "")
33
- messages = data.get("messages", None)
34
- parameters = data.get("parameters", {})
35
-
36
- # 3. 兼容 OpenAI 格式的 Chat completion
37
- if messages is not None:
38
- formatted_input = self.tokenizer.apply_chat_template(
39
- messages, tokenize=False, add_generation_prompt=True
40
- )
41
- else:
42
- formatted_input = inputs
43
 
44
- # 4. 编码输
45
- inputs_tokenized = self.tokenizer([formatted_input], return_tensors="pt").to("cuda")
46
-
47
- # 5. 解析参数
48
- max_new_tokens = parameters.get("max_new_tokens", parameters.get("max_tokens", 128))
49
- temperature = parameters.get("temperature", 0.7)
50
- top_p = parameters.get("top_p", 0.9)
51
- do_sample = parameters.get("do_sample", True)
52
-
53
- # 6. 进行推理
54
- with torch.no_grad():
55
- outputs = self.model.generate(
56
- **inputs_tokenized,
57
- max_new_tokens=max_new_tokens,
58
- temperature=temperature,
59
- top_p=top_p,
60
- do_sample=do_sample
61
- )
62
-
63
- # 7. 只解码新生成的回复部分
64
- input_length = inputs_tokenized.input_ids.shape[1]
65
- reply = self.tokenizer.decode(outputs[0][input_length:], skip_special_tokens=True)
66
-
67
- # 返回符合 Hugging Face 规范的格式
68
- return [{"generated_text": reply}]
 
1
  import subprocess
2
  import sys
3
+ import site
4
 
5
+ # 🚀 1. 强制从 PyPI 热升级 transformers 和 accelerate 到最新版本
6
  try:
7
+ print("Starting hot-upgrade of transformers and accelerate...", flush=True)
8
  subprocess.run(
9
  [sys.executable, "-m", "pip", "install", "--upgrade", "transformers", "accelerate", "--user"],
10
  check=True
11
  )
12
+ print("Hot-upgrade completed successfully!", flush=True)
13
  except Exception as e:
14
+ print(f"Warning during hot-upgrade: {e}", flush=True)
15
 
16
+ # 🚀 2. 强制将 user site-packages sys.path 的最前面,覆盖系统自带的旧版库
17
+ user_site = site.getusersitepackages()
18
+ if user_site in sys.path:
19
+ sys.path.remove(user_site)
20
+ sys.path.insert(0, user_site)
 
 
 
 
 
 
 
 
 
21
 
22
+ # 🚀 3. 清理内存中已加载的旧版模块,强制 Python 重新从新路径导入
23
+ for mod in list(sys.modules.keys()):
24
+ if mod.startswith("transformers") or mod.startswith("accelerate"):
25
+ del sys.modules[mod]
 
 
 
 
 
 
 
 
 
26
 
27
+ # ----------------- 路径重排和模块清理完成后,导新版库 -----------------
28
+ import torch
29
+ from transformers import AutoModelForCausalLM, AutoTokenizer