初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
53
tools/merge_poe_lora.py
Normal file
53
tools/merge_poe_lora.py
Normal file
@@ -0,0 +1,53 @@
|
||||
import argparse
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
from peft import PeftModel
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--base-model", required=True)
|
||||
parser.add_argument("--adapter", required=True)
|
||||
parser.add_argument("--output-dir", required=True)
|
||||
parser.add_argument("--dtype", default="bfloat16", choices=["bfloat16", "float16", "float32"])
|
||||
args = parser.parse_args()
|
||||
|
||||
dtype_map = {
|
||||
"bfloat16": torch.bfloat16,
|
||||
"float16": torch.float16,
|
||||
"float32": torch.float32,
|
||||
}
|
||||
dtype = dtype_map[args.dtype]
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.base_model,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
|
||||
base_model = AutoModelForCausalLM.from_pretrained(
|
||||
args.base_model,
|
||||
torch_dtype=dtype,
|
||||
device_map="auto",
|
||||
trust_remote_code=True,
|
||||
)
|
||||
|
||||
model = PeftModel.from_pretrained(
|
||||
base_model,
|
||||
args.adapter,
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
|
||||
model = model.merge_and_unload()
|
||||
|
||||
model.save_pretrained(
|
||||
args.output_dir,
|
||||
safe_serialization=True,
|
||||
max_shard_size="4GB",
|
||||
)
|
||||
tokenizer.save_pretrained(args.output_dir)
|
||||
|
||||
print(f"Saved merged model to {args.output_dir}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user