refactor: 更新inference 部分的实现
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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()))
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user