feat: 增加推理部分工厂模式
This commit is contained in:
+25
-51
@@ -1,45 +1,10 @@
|
||||
import torch
|
||||
import json
|
||||
import torch
|
||||
import argparse
|
||||
|
||||
from khaosz import Khaosz
|
||||
from typing import List
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def batch_generate(
|
||||
model: Khaosz,
|
||||
query: List[str],
|
||||
temperature: float,
|
||||
top_k: int,
|
||||
top_p: float,
|
||||
batch_size: int,
|
||||
) -> List:
|
||||
assert batch_size > 0
|
||||
sorted_query = sorted(query, key=lambda x: len(x), reverse=True)
|
||||
original_indices = {query: idx for idx, query in enumerate(query)}
|
||||
|
||||
responses = [None] * len(query)
|
||||
total_batches = (len(sorted_query) + batch_size - 1) // batch_size
|
||||
|
||||
for i in tqdm(range(0, total_batches * batch_size, batch_size), desc="Generating responses"):
|
||||
batch_query = sorted_query[i: min(i + batch_size, len(query))]
|
||||
if not isinstance(batch_query, list):
|
||||
batch_query = [batch_query]
|
||||
|
||||
batch_responses = model.batch_generate(
|
||||
query=batch_query,
|
||||
temperature=temperature,
|
||||
top_k=top_k,
|
||||
top_p=top_p
|
||||
)
|
||||
|
||||
for query, response in zip(batch_query, batch_responses):
|
||||
original_idx = original_indices[query]
|
||||
responses[original_idx] = response
|
||||
|
||||
return responses
|
||||
from khaosz.config.param_config import ModelParameter
|
||||
from khaosz.inference.generator import BatchGenerator, GenerationRequest
|
||||
from khaosz.inference.core import disable_random_init
|
||||
|
||||
|
||||
def processor(
|
||||
@@ -53,24 +18,31 @@ def processor(
|
||||
question_key: str,
|
||||
response_key: str,
|
||||
):
|
||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
||||
|
||||
with disable_random_init():
|
||||
param = ModelParameter.load(model_dir)
|
||||
|
||||
param.to(device='cuda', dtype=torch.bfloat16)
|
||||
generator = BatchGenerator(param)
|
||||
|
||||
with open(input_json_file, "r", encoding='utf-8') as f:
|
||||
input_data = [json.loads(line) for line in f]
|
||||
query = [item[question_key] for item in input_data]
|
||||
|
||||
responses = batch_generate(
|
||||
model=model,
|
||||
query=query,
|
||||
|
||||
queries = [item[question_key] for item in input_data]
|
||||
|
||||
request = GenerationRequest(
|
||||
query=queries,
|
||||
temperature=temperature,
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
batch_size=batch_size
|
||||
top_k=top_k,
|
||||
max_len=param.config.max_len,
|
||||
history=None,
|
||||
system_prompt=None,
|
||||
)
|
||||
|
||||
# Write output in JSONL format
|
||||
|
||||
responses = generator.generate(request)
|
||||
|
||||
with open(output_json_file, "w", encoding='utf-8') as f:
|
||||
for query, response in zip(query, responses):
|
||||
for query, response in zip(queries, responses):
|
||||
output_item = {question_key: query, response_key: response}
|
||||
f.write(json.dumps(output_item, ensure_ascii=False) + '\n')
|
||||
|
||||
@@ -89,4 +61,6 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for generating responses.")
|
||||
|
||||
args = parser.parse_args()
|
||||
processor(**vars(args))
|
||||
|
||||
with torch.inference_mode():
|
||||
processor(**vars(args))
|
||||
+13
-15
@@ -6,7 +6,9 @@ import argparse
|
||||
import tqdm
|
||||
|
||||
from torch import Tensor
|
||||
from khaosz import Khaosz
|
||||
from khaosz.config.param_config import ModelParameter
|
||||
from khaosz.inference.core import disable_random_init
|
||||
|
||||
|
||||
def compute_perplexity(
|
||||
model: nn.Module,
|
||||
@@ -45,22 +47,23 @@ def process_file(
|
||||
batch_size: int,
|
||||
text_key: str
|
||||
):
|
||||
model = Khaosz(model_dir).to(device="cuda", dtype=torch.bfloat16)
|
||||
tokenizer = model.parameter.tokenizer
|
||||
with disable_random_init():
|
||||
param = ModelParameter.load(model_dir)
|
||||
|
||||
param.to(device='cuda', dtype=torch.bfloat16)
|
||||
model = param.model
|
||||
tokenizer = param.tokenizer
|
||||
|
||||
with open(input_file, "r", encoding='utf-8') as f:
|
||||
input_data = [json.loads(line) for line in f]
|
||||
|
||||
texts = [item[text_key] for item in input_data]
|
||||
encoded_texts = [tokenizer.encode(text) for text in texts]
|
||||
|
||||
output_data = []
|
||||
|
||||
for i in tqdm(range(0, len(encoded_texts), batch_size), desc="Computing perplexity"):
|
||||
batch_encoded = encoded_texts[i:i + batch_size]
|
||||
batch_texts = texts[i:i + batch_size]
|
||||
|
||||
# Pad sequences to the same length (left padding)
|
||||
max_len = max(len(seq) for seq in batch_encoded)
|
||||
padded_ids = []
|
||||
masks = []
|
||||
@@ -74,10 +77,7 @@ def process_file(
|
||||
|
||||
input_ids = torch.tensor(padded_ids, device="cuda", dtype=torch.long)
|
||||
input_mask = torch.tensor(masks, device="cuda", dtype=torch.bool)
|
||||
|
||||
# Compute perplexity
|
||||
with torch.inference_mode():
|
||||
perplexity = compute_perplexity(model.parameter.model, input_ids, input_mask)
|
||||
perplexity = compute_perplexity(model, input_ids, input_mask)
|
||||
|
||||
for text, ppl in zip(batch_texts, perplexity):
|
||||
output_data.append({text_key: text, "ppl": float(ppl.item())})
|
||||
@@ -87,16 +87,14 @@ def process_file(
|
||||
f.write(json.dumps(item, ensure_ascii=False) + '\n')
|
||||
|
||||
|
||||
def main():
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Run perplexity with a Khaosz model.")
|
||||
parser.add_argument("--model_dir", type=str, required=True, help="Path to the model directory.")
|
||||
parser.add_argument("--input_file", type=str, required=True, help="Path to the input file.")
|
||||
parser.add_argument("--output_file", type=str, required=True, help="Path to the output file.")
|
||||
parser.add_argument("--batch_size", type=int, default=4, help="Batch size for evaluation.")
|
||||
parser.add_argument("--text_key", type=str, default="text", help="Key for the text field in the input data.")
|
||||
|
||||
args = parser.parse_args()
|
||||
process_file(**vars(args))
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
with torch.inference_mode():
|
||||
process_file(**vars(args))
|
||||
|
||||
+4
-9
@@ -16,7 +16,7 @@ def parse_args() -> argparse.Namespace:
|
||||
|
||||
parser = argparse.ArgumentParser(description="Train the Transformer model.")
|
||||
|
||||
parser.add_argument("--train_type",choices=["seq", "sft", "dpo"], help="Train type.")
|
||||
parser.add_argument("--train_type", type=str, required=True, choices=["seq", "sft", "dpo"], help="Train type.")
|
||||
parser.add_argument("--data_root_path", type=str, required=True, help="Path to the root directory of the dataset.")
|
||||
parser.add_argument("--param_path", type=str, required=True, help="Path to the model parameters or resume checkpoint.")
|
||||
|
||||
@@ -67,18 +67,14 @@ def create_scheduler(optimizer: optim.Optimizer, **kwargs) -> optim.lr_scheduler
|
||||
return SchedulerFactory.load(optimizer, **kwargs)
|
||||
|
||||
def prepare_checkpoint(model: nn.Module) -> dict:
|
||||
if isinstance(model, torch.nn.parallel.DistributedDataParallel):
|
||||
state_dict = model.module.state_dict()
|
||||
else:
|
||||
state_dict = model.state_dict()
|
||||
return state_dict
|
||||
return model.module.state_dict()
|
||||
|
||||
|
||||
def train(
|
||||
train_type: str,
|
||||
param_path: str,
|
||||
data_root_path: str,
|
||||
max_lr: int,
|
||||
max_lr: float,
|
||||
n_epoch: int,
|
||||
batch_size: int,
|
||||
start_epoch: int,
|
||||
@@ -104,8 +100,7 @@ def train(
|
||||
assert train_type in ["seq", "sft", "dpo"]
|
||||
assert os.path.exists(param_path)
|
||||
|
||||
parameter = ModelParameter()
|
||||
parameter.load(param_path)
|
||||
parameter = ModelParameter.load(param_path)
|
||||
|
||||
if window_size is None:
|
||||
window_size = parameter.config.max_len
|
||||
|
||||
Reference in New Issue
Block a user