feat: 实现模型动态注册机制

This commit is contained in:
2026-04-05 19:38:12 +08:00
parent ff43a2fab8
commit fc278d17ab
25 changed files with 686 additions and 651 deletions
+13 -5
View File
@@ -2,24 +2,32 @@ from pathlib import Path
import torch
from astrai.config.param_config import ModelParameter
from astrai.inference import InferenceEngine
from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer
PROJECT_ROOT = Path(__file__).resolve().parents[2]
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
def generate_text():
param = ModelParameter.load(PARAMETER_ROOT, disable_init=True)
param.to(device="cuda", dtype=torch.bfloat16)
# Load model from pretrained
model = AutoModel.from_pretrained(PARAMETER_ROOT)
model.to(device="cuda", dtype=torch.bfloat16)
# Load tokenizer from pretrained
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT / "tokenizer")
query = input(">> ")
engine = InferenceEngine(param)
engine = InferenceEngine(
model=model,
tokenizer=tokenizer,
)
response = engine.generate(
prompt=query,
stream=False,
max_tokens=param.config.max_len,
max_tokens=2048,
temperature=0.8,
top_p=0.95,
top_k=50,
+10 -5
View File
@@ -2,7 +2,7 @@ from pathlib import Path
import torch
from astrai.config.param_config import ModelParameter
from astrai.model import AutoModel
from astrai.inference import InferenceEngine
PROJECT_ROOT = Path(__file__).resolve().parents[2]
@@ -10,8 +10,10 @@ PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
def batch_generate():
param = ModelParameter.load(PARAMETER_ROOT, disable_init=True)
param.to(device="cuda", dtype=torch.bfloat16)
# Load model using AutoModel
model = AutoModel.from_pretrained(
PARAMETER_ROOT, device="cuda", dtype=torch.bfloat16
)
inputs = [
"你好",
@@ -21,11 +23,14 @@ def batch_generate():
"请问什么是显卡",
]
engine = InferenceEngine(param)
engine = InferenceEngine(
model=model.model,
tokenizer=model.tokenizer,
)
responses = engine.generate(
prompt=inputs,
stream=False,
max_tokens=param.config.max_len,
max_tokens=model.config.max_len,
temperature=0.8,
top_p=0.95,
top_k=50,
+17 -9
View File
@@ -1,32 +1,39 @@
from pathlib import Path
import torch
from astrai.config.param_config import ModelParameter
from astrai.inference import InferenceEngine
from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer
PROJECT_ROOT = Path(__file__).resolve().parents[2]
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
def chat():
param = ModelParameter.load(PARAMETER_ROOT, disable_init=True)
param.to(device="cuda", dtype=torch.bfloat16)
model = AutoModel.from_pretrained(PARAMETER_ROOT)
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
model.to(device="cuda", dtype=torch.bfloat16)
history = []
engine = InferenceEngine(param)
messages = []
engine = InferenceEngine(model=model, tokenizer=tokenizer)
while True:
query = input(">> ")
if query == "!exit":
break
# Add user message
messages.append({"role": "user", "content": query})
# Generate response
full_response = ""
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
for token in engine.generate(
prompt=query,
prompt=prompt,
stream=True,
max_tokens=param.config.max_len,
max_tokens=model.config.max_len,
temperature=0.8,
top_p=0.95,
top_k=50,
@@ -35,7 +42,8 @@ def chat():
full_response += token
print()
history.append((query, full_response.strip()))
# Add assistant response to messages
messages.append({"role": "assistant", "content": full_response.strip()})
if __name__ == "__main__":