# Set-up the trainer for multilabel finetuning class MultiLabelTrainer(Trainer): def compute_loss(self, model, inputs, return_outputs=False, **kwargs): labels = inputs.pop("labels") outputs = model(**inputs, labels=labels) loss = outputs["loss"] return (loss, outputs) if return_outputs else loss def _save_checkpoint(self, model, trial, metrics=None): super()._save_checkpoint(model, trial) ckpt_dir = self._get_output_dir(trial) # Save head torch.save({ "projection": model.projection.state_dict(), "classifier": model.classifier.state_dict(), }, os.path.join(ckpt_dir, "head_weights.pt")) # Save LoRA adapter explicitly (bypasses bitsandbytes serialization issues) model.backbone.save_pretrained(os.path.join(ckpt_dir, "lora_adapter")) def _load_best_model(self): best_ckpt = self.state.best_model_checkpoint if not best_ckpt: return # Restore head head_path = os.path.join(best_ckpt, "head_weights.pt") if os.path.exists(head_path): head = torch.load(head_path, map_location="cpu") self.model.projection.load_state_dict(head["projection"]) self.model.classifier.load_state_dict(head["classifier"]) print(f"Head restored from: {best_ckpt}") else: print(f"WARNING: head_weights.pt not found in {best_ckpt}") # Restore LoRA adapter lora_path = os.path.join(best_ckpt, "lora_adapter") if os.path.exists(lora_path): from peft import PeftModel self.model.backbone.load_adapter(lora_path, adapter_name="default") print(f"LoRA restored from: {best_ckpt}") else: print(f"WARNING: lora_adapter/ not found in {best_ckpt}") # Launch the trainer trainer = MultiLabelTrainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, compute_metrics=compute_metrics, ) # Launch training trainer.train()