初始化项目,由ModelHub XC社区提供模型

Model: ayh015/myLightningOPD
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-08-27 23:50:14 +08:00
commit d4e0a1af66
368 changed files with 559583 additions and 0 deletions

53
tools/merge_poe_lora.py Normal file
View 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()