style: 使用ruff 工具优化代码风格
This commit is contained in:
+4
-5
@@ -2,13 +2,12 @@ import os
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(
|
||||
os.path.dirname(os.path.abspath(__file__)))
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
snapshot_download(
|
||||
repo_id="ViperEk/KHAOSZ",
|
||||
local_dir=os.path.join(PROJECT_ROOT, "params"),
|
||||
force_download=True
|
||||
)
|
||||
local_dir=os.path.join(PROJECT_ROOT, "params"),
|
||||
force_download=True,
|
||||
)
|
||||
|
||||
+9
-8
@@ -5,18 +5,18 @@ from khaosz.inference.core import disable_random_init
|
||||
from khaosz.inference.generator import LoopGenerator, GenerationRequest
|
||||
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(
|
||||
os.path.dirname(os.path.abspath(__file__)))
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
|
||||
def generate_text():
|
||||
|
||||
|
||||
with disable_random_init():
|
||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
||||
param = ModelParameter.load(model_dir)
|
||||
|
||||
param.to(device='cuda', dtype=torch.bfloat16)
|
||||
|
||||
param.to(device="cuda", dtype=torch.bfloat16)
|
||||
query = input(">> ")
|
||||
|
||||
|
||||
request = GenerationRequest(
|
||||
query=query,
|
||||
temperature=0.8,
|
||||
@@ -28,8 +28,9 @@ def generate_text():
|
||||
)
|
||||
generator = LoopGenerator(param)
|
||||
response = generator.generate(request)
|
||||
|
||||
|
||||
print(response)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
generate_text()
|
||||
generate_text()
|
||||
|
||||
+14
-7
@@ -4,18 +4,24 @@ from khaosz.config.param_config import ModelParameter
|
||||
from khaosz.inference.core import disable_random_init
|
||||
from khaosz.inference.generator import BatchGenerator, GenerationRequest
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(
|
||||
os.path.dirname(os.path.abspath(__file__)))
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
|
||||
def batch_generate():
|
||||
with disable_random_init():
|
||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
||||
param = ModelParameter.load(model_dir)
|
||||
|
||||
param.to(device='cuda', dtype=torch.bfloat16)
|
||||
param.to(device="cuda", dtype=torch.bfloat16)
|
||||
generator = BatchGenerator(param)
|
||||
inputs = ["你好", "请问什么是人工智能", "今天天气如何", "我感到焦虑, 请问我应该怎么办", "请问什么是显卡"]
|
||||
|
||||
inputs = [
|
||||
"你好",
|
||||
"请问什么是人工智能",
|
||||
"今天天气如何",
|
||||
"我感到焦虑, 请问我应该怎么办",
|
||||
"请问什么是显卡",
|
||||
]
|
||||
|
||||
request = GenerationRequest(
|
||||
query=inputs,
|
||||
temperature=0.8,
|
||||
@@ -26,9 +32,10 @@ def batch_generate():
|
||||
system_prompt=None,
|
||||
)
|
||||
responses = generator.generate(request)
|
||||
|
||||
|
||||
for q, r in zip(inputs, responses):
|
||||
print((q, r))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
batch_generate()
|
||||
batch_generate()
|
||||
|
||||
+8
-8
@@ -5,16 +5,16 @@ from khaosz.inference.core import disable_random_init
|
||||
from khaosz.inference.generator import StreamGenerator, GenerationRequest
|
||||
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(
|
||||
os.path.dirname(os.path.abspath(__file__)))
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
|
||||
def chat():
|
||||
|
||||
|
||||
with disable_random_init():
|
||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
||||
param = ModelParameter.load(model_dir)
|
||||
|
||||
param.to(device='cuda', dtype=torch.bfloat16)
|
||||
param.to(device="cuda", dtype=torch.bfloat16)
|
||||
generator = StreamGenerator(param)
|
||||
|
||||
history = []
|
||||
@@ -22,7 +22,7 @@ def chat():
|
||||
query = input(">> ")
|
||||
if query == "!exit":
|
||||
break
|
||||
|
||||
|
||||
request = GenerationRequest(
|
||||
query=query,
|
||||
temperature=0.8,
|
||||
@@ -32,7 +32,7 @@ def chat():
|
||||
history=history,
|
||||
system_prompt=None,
|
||||
)
|
||||
|
||||
|
||||
response_size = 0
|
||||
full_response = ""
|
||||
for response in generator.generate(request):
|
||||
@@ -40,10 +40,10 @@ def chat():
|
||||
print(response[response_size:], end="", flush=True)
|
||||
response_size = len(response)
|
||||
full_response = response
|
||||
|
||||
|
||||
# After generation, update history
|
||||
history.append((query, full_response.strip()))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
chat()
|
||||
chat()
|
||||
|
||||
Reference in New Issue
Block a user