From af6794e424facafe2e390339fd7fce791f84ee59 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 08:18:47 +0000 Subject: [PATCH] add blip2 --- README.md | 6 ++ .../int8_training/fine_tune_blip2_int8.py | 88 +++++++++++++++++++ src/peft/peft_model.py | 2 +- src/peft/tuners/lora.py | 1 - 4 files changed, 95 insertions(+), 2 deletions(-) create mode 100644 examples/int8_training/fine_tune_blip2_int8.py diff --git a/README.md b/README.md index ccdcd55..af84399 100644 --- a/README.md +++ b/README.md @@ -274,6 +274,12 @@ An example is provided in `~examples/causal_language_modeling/peft_lora_clm_acce | ViT | ✅ | | | | | Swin | ✅ | | | | +### Image to text (Multi-modal models) + +| Model | LoRA | Prefix Tuning | P-Tuning | Prompt Tuning | +| --------- | ---- | ---- | ---- | ---- | +| Blip-2 | ✅ | | | | + ___Note that we have tested LoRA for [ViT](https://huggingface.co/docs/transformers/model_doc/vit) and [Swin](https://huggingface.co/docs/transformers/model_doc/swin) for fine-tuning on image classification. However, it should be possible to use LoRA for any compatible model [provided](https://huggingface.co/models?pipeline_tag=image-classification&sort=downloads&search=vit) by 🤗 Transformers. Check out the respective examples to learn more. If you run into problems, please open an issue.___ diff --git a/examples/int8_training/fine_tune_blip2_int8.py b/examples/int8_training/fine_tune_blip2_int8.py new file mode 100644 index 0000000..526336a --- /dev/null +++ b/examples/int8_training/fine_tune_blip2_int8.py @@ -0,0 +1,88 @@ +import torch +from datasets import load_dataset +from torch.utils.data import DataLoader, Dataset +from transformers import AutoModelForVision2Seq, AutoProcessor + +from peft import LoraConfig, get_peft_model + + +config = LoraConfig( + r=16, + lora_alpha=32, + target_modules=["q_proj", "v_proj"], + lora_dropout=0.05, + bias="none", + task_type="VISION_2_SEQ", +) + +model = AutoModelForVision2Seq.from_pretrained("Salesforce/blip2-opt-2.7b", load_in_8bit=True, device_map={"": 0}) +processor = AutoProcessor.from_pretrained("Salesforce/blip2-opt-2.7b") +model = get_peft_model(model, config) + +model.print_trainable_parameters() + +dataset = load_dataset("ybelkada/football-dataset", split="train") + + +class ImageCaptioningDataset(Dataset): + def __init__(self, dataset, processor): + self.dataset = dataset + self.processor = processor + + def __len__(self): + return len(self.dataset) + + def __getitem__(self, idx): + item = self.dataset[idx] + encoding = self.processor(images=item["image"], padding="max_length", return_tensors="pt") + # remove batch dimension + encoding = {k: v.squeeze() for k, v in encoding.items()} + encoding["text"] = item["text"] + return encoding + + +def collator(batch): + # pad the input_ids and attention_mask + processed_batch = {} + for key in batch[0].keys(): + if key != "text": + processed_batch[key] = torch.stack([example[key] for example in batch]) + else: + text_inputs = processor.tokenizer( + [example["text"] for example in batch], padding=True, return_tensors="pt" + ) + processed_batch["input_ids"] = text_inputs["input_ids"] + processed_batch["attention_mask"] = text_inputs["attention_mask"] + return processed_batch + + +train_dataset = ImageCaptioningDataset(dataset, processor) +train_dataloader = DataLoader(train_dataset, shuffle=True, batch_size=2, collate_fn=collator) + +optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5) + +device = "cuda" if torch.cuda.is_available() else "cpu" +model.to(device) + +model.train() + +for epoch in range(50): + print("Epoch:", epoch) + for idx, batch in enumerate(train_dataloader): + input_ids = batch.pop("input_ids").to(device) + pixel_values = batch.pop("pixel_values").to(device, torch.float16) + + outputs = model(input_ids=input_ids, pixel_values=pixel_values, labels=input_ids) + + loss = outputs.loss + + print("Loss:", loss.item()) + + loss.backward() + + optimizer.step() + optimizer.zero_grad() + + if idx % 10 == 0: + generated_output = model.generate(pixel_values=pixel_values) + print(processor.batch_decode(generated_output, skip_special_tokens=True)) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 79d7464..0305b79 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -1102,4 +1102,4 @@ class PeftModelForVision2Seq(PeftModel): output_hidden_states=output_hidden_states, return_dict=return_dict, **kwargs, - ) \ No newline at end of file + ) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index d4d17ac..1674754 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -394,7 +394,6 @@ class Linear(nn.Linear, LoraLayer): self.lora_B.eval() def forward(self, x: torch.Tensor): - if self.disable_adapters: if self.r > 0 and self.merged: self.weight.data -= (