首页 > 教程攻略 > ai资讯 >自训Transformer模型:识别图像是否由AI生成?

自训Transformer模型:识别图像是否由AI生成?

来源:互联网 时间:2026-08-22 20:32:23

这个问题不是抬杠,是实实在在的技术挑战。好在,Transformer架构的出现给了我们一个很不错的解题思路。

Transformer模型在图像识别中的作用

为什么Transformer能胜任这个任务?主要是三个层面的能力在起作用。

特征学习能力

:Transformer的特征提取能力确实强悍,它能从图像中捕捉到极其细粒度的差异——那些AI生成图里不易察觉的微小破绽,比如光影过渡的逻辑瑕疵、纹理生成的自然度偏差,这些在传统模型下很容易被忽略的信息,Transformer能精准抓住。

上下文理解

:和传统CNN相比,Transformer最大的不同在于对全局信息的敏感度。CNN擅长抓局部特征,但对整个画面的语境关联把握不足。Transformer刚好相反,它在分析细节纹理时,会同时考虑周边像素的相互影响,这种全局视角让判断结果更可靠。

适应性强

:AI图像生成技术迭代很快,去年还行得通的识别方法今年可能就失效了。Transformer的预训练+微调机制恰好能应对这个局面——在大规模数据集上打好基础后,针对新的生成技术做微调,模型就能快速跟上节奏,保持高效的识别能力。

下面用一个实战案例来演示具体怎么做。

数据集

类别 训练集数量 测试集数量
FAKE 50,000 10,000
REAL 50,000 10,000

总计

100,000

20,000

完整步骤

1、导入包

import numpy as np
from datasets import load_dataset
import torch
from transformers import ViTFeatureExtractor
from transformers import TrainingArguments
from transformers import Trainer
import tensorflow as tf
from tensorflow.keras.preprocessing.image import ImageDataGenerator
from transformers import ViTForImageClassification, default_data_collator
from torch.utils.data import DataLoader, Dataset
import os
from PIL import Image

2、图像预处理

数据集是标准的文件夹结构,下面按REAL和FAKE分两个子目录存放。我们自定义一个Dataset类来读取数据,同时完成图像的预处理操作。

from torch.utils.data import Dataset
from PIL import Image
import os
import torch

class CustomImageDataset(Dataset):
    def __init__(self, img_dir, feature_extractor):
        self.img_dir = img_dir
        self.img_labels = []
        self.img_files = []
        self.feature_extractor = feature_extractor
        self.label_mapping = {'REAL': 1, 'FAKE': 0}
        
        for label_dir in ['REAL', 'FAKE']:
            dir_path = os.path.join(img_dir, label_dir)
            files = os.listdir(dir_path)
            for file in files:
                self.img_files.append(os.path.join(dir_path, file))
                self.img_labels.append(self.label_mapping[label_dir])

    def __len__(self):
        return len(self.img_files)

    def __getitem__(self, idx):
        img_path = self.img_files[idx]
        image = Image.open(img_path).convert("RGB")
        label = self.img_labels[idx]
        features = self.feature_extractor(images=image, return_tensors="pt")
        return {"pixel_values": features['pixel_values'].squeeze(), "labels": torch.tensor(label)}

3、加载模型与数据加载

这里用Google开源的ViT基础版本——在ImageNet-21k上预训练好的模型,输入分辨率224×224。直接从本地加载模型和特征提取器,然后配置好设备(有GPU优先用GPU,没有就CPU兜底),最后创建训练和测试的数据加载器。

from transformers import ViTFeatureExtractor, ViTForImageClassification
from torch.utils.data import DataLoader
import torch

model_id = 'google/vit-base-patch16-224-in21k'
model_path = '../model'

feature_extractor = ViTFeatureExtractor.from_pretrained(
    model_path, local_files_only=True
)

device = torch.device('cuda' if torch.cuda.is_a vailable() else 'cpu')

model = ViTForImageClassification.from_pretrained(
    model_path, num_labels=2, local_files_only=True
)
model.to(device)

train_dataset = CustomImageDataset(img_dir='../dataset/train', feature_extractor=feature_extractor)
test_dataset = CustomImageDataset(img_dir='../dataset/test', feature_extractor=feature_extractor)

train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False)

4、查看当前设备

print("当前设备:", device)

5、配置训练参数并初始化Trainer

训练参数这块有几个关键设置:4个训练周期,每个周期结束都做一次评估,学习率取2e-4,批次大小设成4(显存有限的话可以考虑调小)。同时加载最佳模型的策略也开启,避免训练到最后反而过拟合。

from transformers import Trainer, TrainingArguments, default_data_collator

training_args = TrainingArguments(
    per_device_train_batch_size=4,
    evaluation_strategy="epoch",
    num_train_epochs=4,
    sa ve_strategy="epoch",
    logging_steps=10,
    learning_rate=2e-4,
    sa ve_total_limit=2,
    remove_unused_columns=False,
    push_to_hub=False,
    load_best_model_at_end=True,
    output_dir="./outputs",
    use_cpu=False
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=test_dataset,
    data_collator=default_data_collator,
    compute_metrics=None,
)

6、开始训练

trainer.train()

训练过程比较耗时,需要耐心等一等。

7、保存模型

训练结束后,把模型权重和配置一起保存到指定目录,方便后续直接加载使用。

trainer.sa ve_model("./outputs/final model")

8、验证效果

拿一张测试图片来看看模型的实际表现。这里用了一张"潘展乐"的图片做验证。从输出结果来看,模型预测为"真实",说明它没有把这张图误判成AI生成。

from transformers import AutoFeatureExtractor, AutoModelForImageClassification
import os
import torch
from PIL import Image
import matplotlib.pyplot as plt
plt.rcParams['font.sans-serif'] = ['Microsoft YaHei']
plt.rcParams['axes.unicode_minus'] = False

model_path = os.path.abspath("D:\MY\8-m\final model")

try:
    feature_extractor = AutoFeatureExtractor.from_pretrained(model_path, local_files_only=True)
except OSError as e:
    print(f"加载特征提取器时出错: {e}")

try:
    model = AutoModelForImageClassification.from_pretrained(model_path, local_files_only=True)
except OSError as e:
    print(f"载模型时出错: {e}")

device = torch.device("cuda" if torch.cuda.is_a vailable() else "cpu")
model.to(device)

image_path = '潘展乐.png'
image = Image.open(image_path).convert("RGB")
inputs = feature_extractor(images=image, return_tensors="pt")
pixel_values = inputs['pixel_values'].to(device)

model.eval()
with torch.no_grad():
    outputs = model(pixel_values)

logits = outputs.logits
predicted_class_idx = logits.argmax(-1).item()
predicted_label = '真实' if predicted_class_idx == 1 else 'AI生成'

print(f"预测标签: {predicted_label}")

plt.imshow(image)
plt.axis('off')
plt.title(f"预测标签: {predicted_label}")
plt.show()

整体来看,模型的准确率还是很棒的。当然,这不是终点——随着生成技术的进化,识别模型也需要持续迭代。但至少目前,Transformer给了我们一个相当可靠的起点。