MCPcopy Create free account
hub / github.com/ZinYY/TreeLoRA / PeftModelForSequenceClassification

Class PeftModelForSequenceClassification

utils/my_peft/peft_model.py:436–621  ·  view source on GitHub ↗

Peft model for sequence classification tasks. Args: model ([`~transformers.PreTrainedModel`]): Base transformer model. peft_config ([`PeftConfig`]): Peft config. **Attributes**: - **config** ([`~transformers.PretrainedConfig`]) -- The configuration object of th

Source from the content-addressed store, hash-verified

434
435
436class PeftModelForSequenceClassification(PeftModel):
437 """
438 Peft model for sequence classification tasks.
439
440 Args:
441 model ([`~transformers.PreTrainedModel`]): Base transformer model.
442 peft_config ([`PeftConfig`]): Peft config.
443
444 **Attributes**:
445 - **config** ([`~transformers.PretrainedConfig`]) -- The configuration object of the base model.
446 - **cls_layer_name** (`str`) -- The name of the classification layer.
447
448 Example:
449
450 ```py
451 >>> from transformers import AutoModelForSequenceClassification
452 >>> from utils.my_peft import PeftModelForSequenceClassification, get_peft_config
453
454 >>> config = {
455 ... "peft_type": "PREFIX_TUNING",
456 ... "task_type": "SEQ_CLS",
457 ... "inference_mode": False,
458 ... "num_virtual_tokens": 20,
459 ... "token_dim": 768,
460 ... "num_transformer_submodules": 1,
461 ... "num_attention_heads": 12,
462 ... "num_layers": 12,
463 ... "encoder_hidden_size": 768,
464 ... "prefix_projection": False,
465 ... "postprocess_past_key_value_function": None,
466 ... }
467
468 >>> peft_config = get_peft_config(config)
469 >>> model = AutoModelForSequenceClassification.from_pretrained("bert-base-cased")
470 >>> peft_model = PeftModelForSequenceClassification(model, peft_config)
471 >>> peft_model.print_trainable_parameters()
472 trainable params: 370178 || all params: 108680450 || trainable%: 0.3406113979101117
473 ```
474 """
475
476 def __init__(self, model, peft_config: PeftConfig, adapter_name="default"):
477 super().__init__(model, peft_config, adapter_name)
478 if self.modules_to_save is None:
479 self.modules_to_save = {"classifier", "score"}
480 else:
481 self.modules_to_save.update({"classifier", "score"})
482
483 for name, _ in self.base_model.named_children():
484 if any(module_name in name for module_name in self.modules_to_save):
485 self.cls_layer_name = name
486 break
487
488 # to make sure classifier layer is trainable
489 _set_trainable(self, adapter_name)
490
491 def forward(
492 self,
493 input_ids=None,

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected