53 lines
1.3 KiB
Python
53 lines
1.3 KiB
Python
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() |