refactor: 更新inference 部分的实现

This commit is contained in:
2026-04-04 23:49:18 +08:00
parent 99b821ebf5
commit 861d33b1a1
13 changed files with 965 additions and 758 deletions
+6 -8
View File
@@ -3,7 +3,7 @@ from pathlib import Path
import torch
from astrai.config.param_config import ModelParameter
from astrai.inference.generator import GenerationRequest, GeneratorFactory
from astrai.inference import InferenceEngine
PROJECT_ROOT = Path(__file__).resolve().parents[2]
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
@@ -15,17 +15,15 @@ def generate_text():
query = input(">> ")
request = GenerationRequest(
query=query,
engine = InferenceEngine(param)
response = engine.generate(
prompt=query,
stream=False,
max_tokens=param.config.max_len,
temperature=0.8,
top_p=0.95,
top_k=50,
max_len=param.config.max_len,
history=None,
system_prompt=None,
)
generator = GeneratorFactory.create(param, request)
response = generator.generate(request)
print(response)
+6 -8
View File
@@ -3,7 +3,7 @@ from pathlib import Path
import torch
from astrai.config.param_config import ModelParameter
from astrai.inference.generator import GenerationRequest, GeneratorFactory
from astrai.inference import InferenceEngine
PROJECT_ROOT = Path(__file__).resolve().parents[2]
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
@@ -21,17 +21,15 @@ def batch_generate():
"请问什么是显卡",
]
request = GenerationRequest(
query=inputs,
engine = InferenceEngine(param)
responses = engine.generate(
prompt=inputs,
stream=False,
max_tokens=param.config.max_len,
temperature=0.8,
top_p=0.95,
top_k=50,
max_len=param.config.max_len,
history=None,
system_prompt=None,
)
generator = GeneratorFactory.create(param, request)
responses = generator.generate(request)
for q, r in zip(inputs, responses):
print((q, r))
+13 -18
View File
@@ -3,7 +3,7 @@ from pathlib import Path
import torch
from astrai.config.param_config import ModelParameter
from astrai.inference.generator import GenerationRequest, GeneratorFactory
from astrai.inference import InferenceEngine
PROJECT_ROOT = Path(__file__).resolve().parents[2]
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
@@ -14,32 +14,27 @@ def chat():
param.to(device="cuda", dtype=torch.bfloat16)
history = []
engine = InferenceEngine(param)
while True:
query = input(">> ")
if query == "!exit":
break
request = GenerationRequest(
query=query,
full_response = ""
for token in engine.generate(
prompt=query,
stream=True,
max_tokens=param.config.max_len,
temperature=0.8,
top_p=0.95,
top_k=50,
max_len=param.config.max_len,
history=history,
system_prompt=None,
stream=True,
)
generator = GeneratorFactory.create(param, request)
):
print(token, end="", flush=True)
full_response += token
response_size = 0
full_response = ""
for response in generator.generate(request):
# response is the cumulative response up to current token
print(response[response_size:], end="", flush=True)
response_size = len(response)
full_response = response
# After generation, update history
print()
history.append((query, full_response.strip()))
+6 -9
View File
@@ -4,7 +4,7 @@ import json
import torch
from astrai.config.param_config import ModelParameter
from astrai.inference.generator import BatchGenerator, GenerationRequest
from astrai.inference import InferenceEngine
def processor(
@@ -19,25 +19,22 @@ def processor(
):
param = ModelParameter.load(model_dir, disable_init=True)
param.to(device="cuda", dtype=torch.bfloat16)
generator = BatchGenerator(param)
engine = InferenceEngine(param)
with open(input_json_file, "r", encoding="utf-8") as f:
input_data = [json.loads(line) for line in f]
queries = [item[question_key] for item in input_data]
request = GenerationRequest(
query=queries,
responses = engine.generate(
prompt=queries,
stream=False,
max_tokens=param.config.max_len,
temperature=temperature,
top_p=top_p,
top_k=top_k,
max_len=param.config.max_len,
history=None,
system_prompt=None,
)
responses = generator.generate(request)
with open(output_json_file, "w", encoding="utf-8") as f:
for query, response in zip(queries, responses):
output_item = {question_key: query, response_key: response}