KBaba7/llama.cpp
0
1import argparse2import glob3import os4import torch5from safetensors import safe_open6from safetensors.torch import save_file7from typing import Any, ContextManager, cast8 9# Function to determine if file is a SafeTensor file10def is_safetensor_file(file_path):11 return file_path.endswith('.safetensors')12 13 14# Unified loading function15def load_model(file_path):16 if is_safetensor_file(file_path):17 tensors = {}18 with cast(ContextManager[Any], safe_open(file_path, framework="pt", device="cpu")) as f:19 for key in f.keys():20 tensors[key] = f.get_tensor(key).clone()21 # output shape22 print(f"{key} : {tensors[key].shape}")23 return tensors, 'safetensor'24 else:25 return torch.load(file_path, map_location=torch.device('cpu')), 'pytorch'26 27 28# Unified saving function29def save_model(model, file_path, file_type):30 if file_type == 'safetensor':31 # safe_save(model, file_path)32 save_file(model, file_path)33 else:34 torch.save(model, file_path)35 36 37# Adapted function to clean vision tower from checkpoint38def clean_vision_tower_from_checkpoint(checkpoint_path):39 checkpoint, file_type = load_model(checkpoint_path)40 # file_type = 'pytorch'41 model_path = os.path.dirname(checkpoint_path)42 print(f"Searching for vision tower tensors in {checkpoint_path}")43 clip_tensors = [k for k, v in checkpoint.items() if (k.startswith("model.vision_tower") or k.startswith("vit."))]44 45 if len(clip_tensors) > 0:46 print(f"Found {len(clip_tensors)} tensors to extract from {checkpoint_path}")47 # Adapted for file type48 clip_path = os.path.join(model_path, "llava.clip")49 50 if os.path.exists(clip_path):51 print(f"Loading existing llava.clip from {clip_path}")52 existing_clip, _ = load_model(clip_path)53 else:54 print(f"Creating new llava.clip at {clip_path}")55 existing_clip = {}56 # Update existing_clip with new tensors, avoid duplicates57 for name in clip_tensors:58 simple_name = name[name.index('vision_model.'):] if 'vision_model.' in name else name59 print(f"Adding {simple_name} to llava.clip")60 if simple_name not in existing_clip:61 existing_clip[simple_name] = checkpoint[name]62 63 # Save the updated clip tensors back to llava.clip64 save_model(existing_clip, clip_path, 'pytorch')65 66 # Remove the tensors from the original checkpoint67 for name in clip_tensors:68 del checkpoint[name]69 70 checkpoint_path = checkpoint_path71 return True72 return False73 74def find_relevant_checkpoints(checkpoint_paths, newline_criteria, projector):75 newline_checkpoint_path = None76 projector_checkpoint_path = None77 78 for path in checkpoint_paths:79 checkpoint, _ = load_model(path)80 if newline_criteria(checkpoint) and newline_checkpoint_path is None:81 newline_checkpoint_path = path82 if projector(checkpoint):83 projector_checkpoint_path = path84 85 return newline_checkpoint_path, projector_checkpoint_path86 87def newline_criteria(checkpoint):88 return any(k.startswith("model.image_newline") for k in checkpoint.keys())89 90def proj_criteria(checkpoint):91 return any(k.startswith("model.mm_projector") or k.startswith("vision_proj.") for k in checkpoint.keys())92 93 94# Command-line interface setup95ap = argparse.ArgumentParser()96ap.add_argument("-m", "--model", required=True, help="Path to LLaVA v1.5+ model")97ap.add_argument("-C", "--clean-vision-tower", action="store_true", help="Remove any vision tower from the model files")98args = ap.parse_args()99 100if args.clean_vision_tower:101 # Generalized to handle both PyTorch and SafeTensors models102 model_files = sorted(glob.glob(f"{args.model}/*"), key=os.path.getmtime, reverse=True)103 # checkpoint_paths = [path for path in model_files if (path.endswith('.bin') and path.startswith('pytorch')) or (path.endswith('.safetensors') and path.startswith('model'))]104 checkpoint_paths = [path for path in model_files if (path.endswith('.bin') and 'pytorch' in path.split('/')[-1].split('\\')[-1]) or (path.endswith('.safetensors') and 'model' in path.split('/')[-1].split('\\')[-1])]105 for projector_checkpoint_path in checkpoint_paths:106 print(f"Cleaning {projector_checkpoint_path}")107 if not clean_vision_tower_from_checkpoint(projector_checkpoint_path):108 print(f"No vision tower found in {projector_checkpoint_path}")109 # we break once none is found, so far all models append them at the end110 # break111 print("Done! All vision tower tensors are removed from the model files and stored in llava.clip file.")112 113# Now we look for the projector in the last checkpoint114model_files = sorted(glob.glob(f"{args.model}/*"), key=os.path.getmtime, reverse=True)115checkpoint_paths = [path for path in model_files if (path.endswith('.bin') and 'pytorch' in path.split('/')[-1].split('\\')[-1]) or (path.endswith('.safetensors') and 'model' in path.split('/')[-1].split('\\')[-1])]116# last_checkpoint_path = checkpoint_paths[0]117# first_checkpoint_path = checkpoint_paths[-1]118newline_checkpoint_path, projector_checkpoint_path = find_relevant_checkpoints(checkpoint_paths, newline_criteria, proj_criteria)119 120print(f"Taking projector from {projector_checkpoint_path}")121first_mm_tensors = []122first_checkpoint = None123if newline_checkpoint_path is not None:124 print(f"Taking newline from {newline_checkpoint_path}")125 first_checkpoint, file_type = load_model(newline_checkpoint_path)126 first_mm_tensors = [k for k, v in first_checkpoint.items() if k.startswith("model.image_newline")]127 128# Load the checkpoint129mm_tensors = []130last_checkpoint = None131if projector_checkpoint_path is not None:132 last_checkpoint, file_type = load_model(projector_checkpoint_path)133 mm_tensors = [k for k, v in last_checkpoint.items() if k.startswith("model.mm_projector") or k.startswith("vision_proj.")]134 135if len(mm_tensors) == 0:136 if last_checkpoint is not None:137 for k, v in last_checkpoint.items():138 print(k)139 print(f"Found {len(mm_tensors)} tensors to extract out of {len(last_checkpoint) if last_checkpoint is not None else 0} tensors.")140 print("No tensors found. Is this a LLaVA model?")141 exit()142 143print(f"Found {len(mm_tensors)} tensors to extract.")144print(f"Found additional {len(first_mm_tensors)} tensors to extract.")145# projector = {name: checkpoint.[name].float() for name in mm_tensors}146projector = {}147for name in mm_tensors:148 assert last_checkpoint is not None149 projector[name] = last_checkpoint[name].float()150for name in first_mm_tensors:151 assert first_checkpoint is not None152 projector[name] = first_checkpoint[name].float()153 154if len(projector) > 0:155 save_model(projector, f"{args.model}/llava.projector", 'pytorch')156 157print("Done!")158print(f"Now you can convert {args.model} to a a regular LLaMA GGUF file.")159print(f"Also, use {args.model}/llava.projector to prepare a llava-encoder.gguf file.")160 