style: 使用ruff 工具优化代码风格

This commit is contained in:
2026-03-30 23:32:28 +08:00
parent 345fd2f091
commit 426af2d75f
52 changed files with 1836 additions and 1493 deletions
+4 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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()