# pip install -U bitsandbytes>=0.46.1 import torch import torch.nn as nn from transformers import ( Blip2Processor, Blip2ForConditionalGeneration, BitsAndBytesConfig ) from peft import LoraConfig, get_peft_model # ========================== # LOAD BLIP2 # ========================== quant_config = BitsAndBytesConfig( load_in_8bit=True ) processor = Blip2Processor.from_pretrained( "Salesforce/blip2-opt-2.7b" ) base_model = Blip2ForConditionalGeneration.from_pretrained( "Salesforce/blip2-opt-2.7b", quantization_config=quant_config, device_map="auto" ) # ========================== # RECREATE LORA # ========================== lora_config = LoraConfig( r=32, lora_alpha=64, target_modules=["q_proj", "k_proj", "v_proj", "o_proj"], lora_dropout=0.1, bias="none" ) model = get_peft_model( base_model, lora_config ) # ========================== # RECREATE HYBRID MODEL # ========================== class AgriVisionHybridModel(nn.Module): def __init__(self, blip2_model, num_disease, num_pathogen): super().__init__() self.blip2 = blip2_model hidden_size = self.blip2.language_model.config.hidden_size lm_device = next( self.blip2.language_model.parameters() ).device self.disease_classifier = nn.Linear( hidden_size, num_disease ).to(lm_device) self.pathogen_classifier = nn.Linear( hidden_size, num_pathogen ).to(lm_device) self.ce_loss = nn.CrossEntropyLoss() self.lm_device = lm_device # Your notebook values NUM_DISEASES = 17 NUM_PATHOGENS = 6 hybrid_model = AgriVisionHybridModel( model, num_disease=NUM_DISEASES, num_pathogen=NUM_PATHOGENS ) print("Hybrid model rebuilt") checkpoint = torch.load( "/kaggle/input/models/anhadmahajan06/agrivision-blip2-model/pytorch/default/1/agrivision_final_model.pth", map_location="cpu" ) hybrid_model.load_state_dict( checkpoint["model_state_dict"], strict=False ) print("Checkpoint loaded") hybrid_model.blip2.save_pretrained( "agrivision_adapter" ) processor.save_pretrained( "agrivision_adapter" ) import os for f in os.listdir("agrivision_adapter"): print(f) torch.save( { "disease_classifier": hybrid_model.disease_classifier.state_dict(), "pathogen_classifier": hybrid_model.pathogen_classifier.state_dict() }, "classification_heads.pth" ) get_ipython().getoutput("pip install -q huggingface_hub") from huggingface_hub import login login() from huggingface_hub import upload_folder # Push your model files upload_folder(folder_path=".", repo_id="AnhadMahajan/AgriVision-BLIP2", repo_type="model")