前置准备
环境依赖安装
请按以下命令安装所需工具和库,确保每一步无报错:
- 安装Python3.8+:直接从Python官网下载对应系统版本,安装时勾选
Add Python to PATH
- 执行安装命令:
pip install torch torchvision pandas opencv-python flask
核心实操步骤
第一步:准备档案数据
1. 收集1000张以上带分类标签的档案图片,分类定义为:0=合同、1=报表、2=证件(可从公开数据集或自有档案中提取);
2. 所有图片必须按type_序号.jpg格式命名,例如合同类第一张为type_0_001.jpg,报表类第十张为type_1_010.jpg;
3. 执行以下Python代码生成标签文件,代码完整可直接复制:
```python
import os
import pandas as pd
image_dir = "./archive_images"
labels = []
for img_name in os.listdir(image_dir):
if img_name.endswith(".jpg"):
label = int(img_name.split("_")[1])
labels.append({"image": img_name, "label": label})
df = pd.DataFrame(labels)
df.to_csv("archive_labels.csv", index=False)
```
⚠️ 重点:必须确保图片目录中只有分类.jpg文件,否则代码会报错。
第二步:搭建轻量化分类模型
本步骤使用预训练MobileNetV2模型,无需从零训练,降低门槛;

执行以下代码创建模型脚本train_model.py:
```python
import torch
import torch.nn as nn
from torchvision import models, transforms
from torch.utils.data import Dataset, DataLoader
from sklearn.model_selection import train_test_split
import pandas as pd
from PIL import Image
配置参数
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
batch_size = 16
num_epochs = 10
num_classes = 3
数据预处理
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
自定义数据集类
class ArchiveDataset(Dataset):
def __init__(self, df, img_dir, transform=None):
self.df = df
self.img_dir = img_dir
self.transform = transform
def __len__(self):
return len(self.df)
def __getitem__(self, idx):
img_name = os.path.join(self.img_dir, self.df.iloc[idx, 0])
image = Image.open(img_name).convert("RGB")
label = self.df.iloc[idx, 1]
if self.transform:
image = self.transform(image)
return image, label
加载数据
df = pd.read_csv("archive_labels.csv")
train_df, val_df = train_test_split(df, test_size=0.2, random_state=42)
train_dataset = ArchiveDataset(train_df, "./archive_images", transform)
val_dataset = ArchiveDataset(val_df, "./archive_images", transform)
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)
加载并修改预训练模型
model = models.mobilenet_v2(pretrained=True)
model.classifier[1] = nn.Linear(model.classifier[1].in_features, num_classes)
model = model.to(device)
损失函数和优化器
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
```
第三步:模型训练与验证
在train_model.py中添加以下训练代码,完整脚本如下:
```python
训练循环
for epoch in range(num_epochs):
model.train()
train_loss = 0.0
for images, labels in train_loader:
images, labels = images.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
train_loss += loss.item() images.size(0)
train_loss /= len(train_loader.dataset)
验证
model.eval()
val_loss = 0.0
correct = 0
total = 0
with torch.no_grad():
for images, labels in val_loader:
images, labels = images.to(device), labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
val_loss += loss.item() images.size(0)
_, predicted = outputs.max(1)
total += labels.size(0)
correct += predicted.eq(labels).sum().item()
val_loss /= len(val_loader.dataset)
val_acc = 100 correct / total
打印进度
print(f"Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%")
保存模型
torch.save(model.state_dict(), "archive_model.pth")
```
⚠️ 重点:当验证准确率达到85%以上时,停止训练;若连续2个epoch准确率无提升,需调整lr至5e-5。
第四步:部署到档案管理软件
1. 创建推理脚本infer_archive.py,用于接收档案图片并输出分类结果,代码完整可复制:
```python
import torch
from torchvision import transforms
from PIL import Image
from train_model import model, device
加载模型权重
model.load_state_dict(torch.load("archive_model.pth"))
model.eval()
预处理
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
推理函数
def infer_archive(img_path):
image = Image.open(img_path).convert("RGB")
image = transform(image).unsqueeze(0).to(device)
with torch.no_grad():
outputs = model(image)
_, predicted = outputs.max(1)
分类映射:0=合同,1=报表,2=证件
class_map = {0: "合同", 1: "报表", 2: "证件"}
return class_map[predicted.item()]
测试:执行python infer_archive.py 档案图片路径
if __name__ == "__main__":
import sys
img_path = sys.argv[1]
result = infer_archive(img_path)
print(f"档案类型:{result}")
```
2. 部署步骤:将archive_model.pth、infer_archive.py与你的档案管理软件放在同一服务器目录,在软件的档案上传接口中调用推理脚本,传入图片路径即可自动识别档案类型;
3. 快速验证:执行python infer_archive.py 你的档案图片路径,直接查看输出结果。
常见问题排查
- 报错“CUDA out of memory”:减少batch_size至8,或使用CPU训练(将device改为
torch.device("cpu"));
- 模型准确率低:检查图片命名是否符合规范,或增加训练数据至2000张以上;
- 依赖安装失败:执行
pip install --upgrade pip后重新安装所有库。