You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

微调LayoutLMv3分类模型时的多类别扩展问题求助

解决LayoutLMv3多类别微调时复用原分类层权重的问题

当基于已训练的4类LayoutLMv3文档分类模型,微调新增类别的数据集时,直接使用ignore_mismatched_sizes=True加载模型会导致分类器层被重新初始化,无法复用原模型对4类的识别权重。以下是具体修改方案:

核心思路

手动加载原模型的编码器权重和分类器中对应原类别的权重,仅对新增类别的分类器权重进行随机初始化,同时调整微调策略保护原有特征提取能力。

具体修改步骤

1. 替换模型初始化逻辑,手动复用权重

不再直接通过from_pretrained加载原模型并修改num_labels,而是分步骤复制权重:

from transformers import LayoutLMv3ForSequenceClassification

# 加载已训练的4类模型
original_model = LayoutLMv3ForSequenceClassification.from_pretrained('user/mymodel1')
# 初始化支持7类的新模型(基于base架构)
new_model = LayoutLMv3ForSequenceClassification.from_pretrained(
    "microsoft/layoutlmv3-base",
    num_labels=7
)

# 复制原模型的编码器权重(保留LayoutLMv3的特征提取能力)
new_model.layoutlmv3.load_state_dict(original_model.layoutlmv3.state_dict())

# 复制原分类器中对应4个类别的权重到新分类器
# 处理输出层权重:原形状[4,768] → 新形状[7,768],前4行复用原权重
new_model.classifier.out_proj.weight.data[:4] = original_model.classifier.out_proj.weight.data
# 处理输出层偏置:原形状[4] → 新形状[7],前4个值复用原偏置
new_model.classifier.out_proj.bias.data[:4] = original_model.classifier.out_proj.bias.data

# 更新模型的label映射(替换为你的新旧类别组合)
new_model.config.id2label = {
    0: "原类别1", 1: "原类别2", 2: "原类别3", 3: "原类别4",
    4: "新类别1", 5: "新类别2", 6: "新类别3"
}
new_model.config.label2id = {v: k for k, v in new_model.config.id2label.items()}

2. 修改PyTorch Lightning的ModelModule类

将上述权重复用逻辑整合到Lightning模块的初始化中,同时调整优化器策略:

class ModelModule(pl.LightningModule):
    def __init__(self, original_model_path: str, new_label_map: dict):
        super().__init__()
        num_labels = len(new_label_map)
        # 加载原模型并初始化新模型
        original_model = LayoutLMv3ForSequenceClassification.from_pretrained(original_model_path)
        self.model = LayoutLMv3ForSequenceClassification.from_pretrained(
            "microsoft/layoutlmv3-base",
            num_labels=num_labels
        )
        # 复用编码器权重
        self.model.layoutlmv3.load_state_dict(original_model.layoutlmv3.state_dict())
        # 复用原类别对应的分类器权重
        self.model.classifier.out_proj.weight.data[:4] = original_model.classifier.out_proj.weight.data
        self.model.classifier.out_proj.bias.data[:4] = original_model.classifier.out_proj.bias.data
        # 更新label映射
        self.model.config.id2label = new_label_map
        self.model.config.label2id = {v: k for k, v in new_label_map.items()}
        # 初始化准确率指标
        self.train_accuracy = Accuracy(task="multiclass", num_classes=num_labels)
        self.val_accuracy = Accuracy(task="multiclass", num_classes=num_labels)

    def forward(self, input_ids, attention_mask, bbox, pixel_values, labels=None):
        return self.model(
            input_ids,
            attention_mask=attention_mask,
            bbox=bbox,
            pixel_values=pixel_values,
            labels=labels
        )

    def training_step(self, batch, batch_idx):
        input_ids = batch["input_ids"]
        attention_mask = batch["attention_mask"]
        bbox = batch["bbox"]
        pixel_values = batch["pixel_values"]
        labels = batch["labels"]
        output = self(input_ids, attention_mask, bbox, pixel_values, labels)
        self.log("train_loss", output.loss)
        self.log(
            "train_acc",
            self.train_accuracy(output.logits, labels),
            on_step=True,
            on_epoch=True
        )
        return output.loss

    def validation_step(self, batch, batch_idx):
        input_ids = batch["input_ids"]
        attention_mask = batch["attention_mask"]
        bbox = batch["bbox"]
        pixel_values = batch["pixel_values"]
        labels = batch["labels"]
        output = self(input_ids, attention_mask, bbox, pixel_values, labels)
        self.log("val_loss", output.loss)
        self.log(
            "val_acc",
            self.val_accuracy(output.logits, labels),
            on_step=False,
            on_epoch=True
        )
        return output.loss

    def configure_optimizers(self):
        # 分层设置学习率:编码器用低学习率,分类器用高学习率
        param_groups = [
            {"params": self.model.layoutlmv3.parameters(), "lr": 1e-5},
            {"params": self.model.classifier.parameters(), "lr": 1e-4}
        ]
        optimizer = torch.optim.Adam(param_groups)
        return optimizer

3. 调整微调策略

  • 初始阶段冻结编码器:先只训练分类器层(尤其是新增类别的部分),训练2-3轮后再解冻编码器,避免原有的特征提取能力被破坏。
  • 分层学习率:如上述代码所示,给编码器设置较低的学习率(1e-5),分类器设置较高的学习率(1e-4),平衡原有权重的保留和新类别的适配。

为什么原方法失效

使用ignore_mismatched_sizes=True时,Hugging Face的模型加载器会跳过所有形状不匹配的层,直接重新初始化。这意味着原模型中针对4类训练的分类器权重完全被丢弃,新模型的分类器全为随机初始化,无法复用原有的识别能力。

内容的提问来源于stack exchange,提问作者im_s_kumar

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 11:50:58