From c2ef46f1454987f2fed2b7d45b655c75498022a5 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 07:58:48 +0000 Subject: [PATCH 1/9] v1 --- src/peft/peft_model.py | 69 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 69 insertions(+) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index f9573bb..fbb4867 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -1034,3 +1034,72 @@ class PeftModelForTokenClassification(PeftModel): hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) + + +class PeftModelForVision2Seq(PeftModel): + """ + Peft model for vision to text models. + + Args: + model ([`~transformers.PreTrainedModel`]): Base transformer model. + peft_config ([`PeftConfig`]): Peft config. + + + Example: + + ```py + >>> from transformers import AutoModelForVision2Seq + >>> from peft import PeftModelForVision2Seq, get_peft_config + + >>> config = { + ... "peft_type": "LORA", + ... "task_type": "VISION_2_SEQ", + ... "inference_mode": False, + ... "r": 8, + ... "target_modules": ["q", "v"], + ... "lora_alpha": 32, + ... "lora_dropout": 0.1, + ... "merge_weights": False, + ... "fan_in_fan_out": False, + ... "enable_lora": None, + ... "bias": "none", + ... } + + >>> peft_config = get_peft_config(config) + >>> model = AutoModelForCausalLM.from_pretrained("Salesforce/blip2-flan-t5-xl") + >>> peft_model = PeftModelForVision2Seq(model, peft_config) + >>> peft_model.print_trainable_parameters() + trainable params: 1843200 || all params: 775873280 || trainable%: 0.23756456724479544 + ``` + """ + + def __init__(self, model, peft_config: PeftConfig): + super().__init__(model, peft_config) + self.base_model_prepare_inputs_for_generation = self.base_model.prepare_inputs_for_generation + + def forward( + self, + pixel_values=None, + attention_mask=None, + decoder_input_ids=None, + decoder_attention_mask=None, + labels=None, + output_attentions=None, + output_hidden_states=None, + return_dict=None, + **kwargs, + ): + r""" + A simple wrapper around the base model's forward method. + """ + return self.base_model( + pixel_values=pixel_values, + attention_mask=attention_mask, + decoder_input_ids=decoder_input_ids, + decoder_attention_mask=decoder_attention_mask, + labels=labels, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + **kwargs, + ) \ No newline at end of file From c7e22ccd757c7f5ba5e459bcf86416fd803c3afa Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 07:59:03 +0000 Subject: [PATCH 2/9] v1 --- src/peft/mapping.py | 8 +++++++- src/peft/tuners/lora.py | 11 +++++++++-- 2 files changed, 16 insertions(+), 3 deletions(-) diff --git a/src/peft/mapping.py b/src/peft/mapping.py index dbb9f36..1e98edb 100644 --- a/src/peft/mapping.py +++ b/src/peft/mapping.py @@ -19,6 +19,7 @@ from .peft_model import ( PeftModelForSeq2SeqLM, PeftModelForSequenceClassification, PeftModelForTokenClassification, + PeftModelForVision2Seq, ) from .tuners import LoraConfig, PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig from .utils import PromptLearningConfig @@ -29,6 +30,7 @@ MODEL_TYPE_TO_PEFT_MODEL_MAPPING = { "SEQ_2_SEQ_LM": PeftModelForSeq2SeqLM, "CAUSAL_LM": PeftModelForCausalLM, "TOKEN_CLS": PeftModelForTokenClassification, + "VISION_2_SEQ": PeftModelForVision2Seq, } PEFT_TYPE_TO_CONFIG_MAPPING = { @@ -44,6 +46,7 @@ TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING = { "bart": ["q_proj", "v_proj"], "gpt2": ["c_attn"], "bloom": ["query_key_value"], + "blip2": ["q", "v", "q_proj", "v_proj"], "opt": ["q_proj", "v_proj"], "gptj": ["q_proj", "v_proj"], "gpt_neox": ["query_key_value"], @@ -134,9 +137,12 @@ def get_peft_model(model, peft_config): model ([`transformers.PreTrainedModel`]): Model to be wrapped. peft_config ([`PeftConfig`]): Configuration object containing the parameters of the Peft model. """ - model_config = model.config.to_dict() peft_config.base_model_name_or_path = model.__dict__.get("name_or_path", None) + + if peft_config.task_type == "VISION_2_SEQ" and not isinstance(peft_config, LoraConfig): + raise ValueError("Vision2Seq task type is only supported with LORA") + if peft_config.task_type not in MODEL_TYPE_TO_PEFT_MODEL_MAPPING.keys(): peft_config = _prepare_lora_config(peft_config, model_config) return PeftModel(model, peft_config) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 51cd56f..d4d17ac 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -394,6 +394,7 @@ 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 -= ( @@ -401,14 +402,20 @@ class Linear(nn.Linear, LoraLayer): ) self.merged = False - return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + + return result elif self.r > 0 and not self.merged: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) if self.r > 0: + x = x.to(self.lora_A.weight.dtype) + result += self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling return result else: - return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + + return result class MergedLinear(nn.Linear, LoraLayer): From af6794e424facafe2e390339fd7fce791f84ee59 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 08:18:47 +0000 Subject: [PATCH 3/9] 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 -= ( From f569bc682bb1998598f645ca9b5c98f847f5e90c Mon Sep 17 00:00:00 2001 From: Younes Belkada <49240599+younesbelkada@users.noreply.github.com> Date: Tue, 4 Apr 2023 10:21:38 +0200 Subject: [PATCH 4/9] Update src/peft/peft_model.py --- src/peft/peft_model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 0305b79..6706b23 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -1066,7 +1066,7 @@ class PeftModelForVision2Seq(PeftModel): ... } >>> peft_config = get_peft_config(config) - >>> model = AutoModelForCausalLM.from_pretrained("Salesforce/blip2-flan-t5-xl") + >>> model = AutoModelForVision2Seq.from_pretrained("Salesforce/blip2-flan-t5-xl") >>> peft_model = PeftModelForVision2Seq(model, peft_config) >>> peft_model.print_trainable_parameters() trainable params: 1843200 || all params: 775873280 || trainable%: 0.23756456724479544 From 46ab59628cd75158ef232e9a63999626a9d9d949 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 08:23:47 +0000 Subject: [PATCH 5/9] revert --- src/peft/tuners/lora.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 1674754..0553f18 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -412,9 +412,8 @@ class Linear(nn.Linear, LoraLayer): result += self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling return result else: - result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) + return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) - return result class MergedLinear(nn.Linear, LoraLayer): From 96cd0390367f73ba494e23e6cd51b75b3df17a16 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 08:29:51 +0000 Subject: [PATCH 6/9] fix --- examples/int8_training/fine_tune_blip2_int8.py | 1 - src/peft/mapping.py | 2 +- src/peft/tuners/lora.py | 1 - 3 files changed, 1 insertion(+), 3 deletions(-) diff --git a/examples/int8_training/fine_tune_blip2_int8.py b/examples/int8_training/fine_tune_blip2_int8.py index 526336a..1e73ffb 100644 --- a/examples/int8_training/fine_tune_blip2_int8.py +++ b/examples/int8_training/fine_tune_blip2_int8.py @@ -9,7 +9,6 @@ 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", diff --git a/src/peft/mapping.py b/src/peft/mapping.py index 1e98edb..4abbd5a 100644 --- a/src/peft/mapping.py +++ b/src/peft/mapping.py @@ -46,7 +46,7 @@ TRANSFORMERS_MODELS_TO_LORA_TARGET_MODULES_MAPPING = { "bart": ["q_proj", "v_proj"], "gpt2": ["c_attn"], "bloom": ["query_key_value"], - "blip2": ["q", "v", "q_proj", "v_proj"], + "blip-2": ["q", "v", "q_proj", "v_proj"], "opt": ["q_proj", "v_proj"], "gptj": ["q_proj", "v_proj"], "gpt_neox": ["query_key_value"], diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 0553f18..6fae36e 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -415,7 +415,6 @@ class Linear(nn.Linear, LoraLayer): return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) - class MergedLinear(nn.Linear, LoraLayer): # Lora implemented in a dense layer def __init__( From 4cbd6cfd43c76d4762dcf93ed97cfa803b61a84f Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 08:31:37 +0000 Subject: [PATCH 7/9] revert --- src/peft/tuners/lora.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 6fae36e..9745467 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -401,9 +401,7 @@ class Linear(nn.Linear, LoraLayer): ) self.merged = False - result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) - - return result + return F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) elif self.r > 0 and not self.merged: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) if self.r > 0: From 8c83386ef413174bb2ee7167d1d13fa0fe00516c Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 08:37:32 +0000 Subject: [PATCH 8/9] few fixes --- .../int8_training/fine_tune_blip2_int8.py | 21 +++++++++++++++++-- src/peft/tuners/lora.py | 2 -- 2 files changed, 19 insertions(+), 4 deletions(-) diff --git a/examples/int8_training/fine_tune_blip2_int8.py b/examples/int8_training/fine_tune_blip2_int8.py index 1e73ffb..25121f6 100644 --- a/examples/int8_training/fine_tune_blip2_int8.py +++ b/examples/int8_training/fine_tune_blip2_int8.py @@ -1,3 +1,17 @@ +# coding=utf-8 +# Copyright 2023-present the HuggingFace Inc. team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. import torch from datasets import load_dataset from torch.utils.data import DataLoader, Dataset @@ -6,6 +20,7 @@ from transformers import AutoModelForVision2Seq, AutoProcessor from peft import LoraConfig, get_peft_model +# Let's define the LoraConfig config = LoraConfig( r=16, lora_alpha=32, @@ -14,12 +29,15 @@ config = LoraConfig( task_type="VISION_2_SEQ", ) +# We load our model and processor using `transformers` 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) +# Get our peft model and print the number of trainable parameters +model = get_peft_model(model, config) model.print_trainable_parameters() +# Let's load the dataset here! dataset = load_dataset("ybelkada/football-dataset", split="train") @@ -61,7 +79,6 @@ train_dataloader = DataLoader(train_dataset, shuffle=True, batch_size=2, collate optimizer = torch.optim.AdamW(model.parameters(), lr=5e-5) device = "cuda" if torch.cuda.is_available() else "cpu" -model.to(device) model.train() diff --git a/src/peft/tuners/lora.py b/src/peft/tuners/lora.py index 9745467..51cd56f 100644 --- a/src/peft/tuners/lora.py +++ b/src/peft/tuners/lora.py @@ -405,8 +405,6 @@ class Linear(nn.Linear, LoraLayer): elif self.r > 0 and not self.merged: result = F.linear(x, transpose(self.weight, self.fan_in_fan_out), bias=self.bias) if self.r > 0: - x = x.to(self.lora_A.weight.dtype) - result += self.lora_B(self.lora_A(self.lora_dropout(x))) * self.scaling return result else: From 7ed9ad04bfa775397e9cb4f112f8d6f04df6c4e0 Mon Sep 17 00:00:00 2001 From: younesbelkada Date: Tue, 4 Apr 2023 11:01:07 +0000 Subject: [PATCH 9/9] revert changes --- .../int8_training/fine_tune_blip2_int8.py | 1 - src/peft/mapping.py | 5 -- src/peft/peft_model.py | 69 ------------------- 3 files changed, 75 deletions(-) diff --git a/examples/int8_training/fine_tune_blip2_int8.py b/examples/int8_training/fine_tune_blip2_int8.py index 25121f6..ca6ba40 100644 --- a/examples/int8_training/fine_tune_blip2_int8.py +++ b/examples/int8_training/fine_tune_blip2_int8.py @@ -26,7 +26,6 @@ config = LoraConfig( lora_alpha=32, lora_dropout=0.05, bias="none", - task_type="VISION_2_SEQ", ) # We load our model and processor using `transformers` diff --git a/src/peft/mapping.py b/src/peft/mapping.py index 4abbd5a..35a6901 100644 --- a/src/peft/mapping.py +++ b/src/peft/mapping.py @@ -19,7 +19,6 @@ from .peft_model import ( PeftModelForSeq2SeqLM, PeftModelForSequenceClassification, PeftModelForTokenClassification, - PeftModelForVision2Seq, ) from .tuners import LoraConfig, PrefixTuningConfig, PromptEncoderConfig, PromptTuningConfig from .utils import PromptLearningConfig @@ -30,7 +29,6 @@ MODEL_TYPE_TO_PEFT_MODEL_MAPPING = { "SEQ_2_SEQ_LM": PeftModelForSeq2SeqLM, "CAUSAL_LM": PeftModelForCausalLM, "TOKEN_CLS": PeftModelForTokenClassification, - "VISION_2_SEQ": PeftModelForVision2Seq, } PEFT_TYPE_TO_CONFIG_MAPPING = { @@ -140,9 +138,6 @@ def get_peft_model(model, peft_config): model_config = model.config.to_dict() peft_config.base_model_name_or_path = model.__dict__.get("name_or_path", None) - if peft_config.task_type == "VISION_2_SEQ" and not isinstance(peft_config, LoraConfig): - raise ValueError("Vision2Seq task type is only supported with LORA") - if peft_config.task_type not in MODEL_TYPE_TO_PEFT_MODEL_MAPPING.keys(): peft_config = _prepare_lora_config(peft_config, model_config) return PeftModel(model, peft_config) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index 6706b23..85757b7 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -1034,72 +1034,3 @@ class PeftModelForTokenClassification(PeftModel): hidden_states=outputs.hidden_states, attentions=outputs.attentions, ) - - -class PeftModelForVision2Seq(PeftModel): - """ - Peft model for vision to text models. - - Args: - model ([`~transformers.PreTrainedModel`]): Base transformer model. - peft_config ([`PeftConfig`]): Peft config. - - - Example: - - ```py - >>> from transformers import AutoModelForVision2Seq - >>> from peft import PeftModelForVision2Seq, get_peft_config - - >>> config = { - ... "peft_type": "LORA", - ... "task_type": "VISION_2_SEQ", - ... "inference_mode": False, - ... "r": 8, - ... "target_modules": ["q", "v"], - ... "lora_alpha": 32, - ... "lora_dropout": 0.1, - ... "merge_weights": False, - ... "fan_in_fan_out": False, - ... "enable_lora": None, - ... "bias": "none", - ... } - - >>> peft_config = get_peft_config(config) - >>> model = AutoModelForVision2Seq.from_pretrained("Salesforce/blip2-flan-t5-xl") - >>> peft_model = PeftModelForVision2Seq(model, peft_config) - >>> peft_model.print_trainable_parameters() - trainable params: 1843200 || all params: 775873280 || trainable%: 0.23756456724479544 - ``` - """ - - def __init__(self, model, peft_config: PeftConfig): - super().__init__(model, peft_config) - self.base_model_prepare_inputs_for_generation = self.base_model.prepare_inputs_for_generation - - def forward( - self, - pixel_values=None, - attention_mask=None, - decoder_input_ids=None, - decoder_attention_mask=None, - labels=None, - output_attentions=None, - output_hidden_states=None, - return_dict=None, - **kwargs, - ): - r""" - A simple wrapper around the base model's forward method. - """ - return self.base_model( - pixel_values=pixel_values, - attention_mask=attention_mask, - decoder_input_ids=decoder_input_ids, - decoder_attention_mask=decoder_attention_mask, - labels=labels, - output_attentions=output_attentions, - output_hidden_states=output_hidden_states, - return_dict=return_dict, - **kwargs, - )