Проблема обучения PyTorch с EfficientNetB3 | Табличка точности проверкиPython

Программы на Python
Anonymous
Проблема обучения PyTorch с EfficientNetB3 | Табличка точности проверки

Сообщение Anonymous »

Метрики
Прежде чем начать, я хотел бы добавить, что я относительно новичок в глубоком обучении, поэтому моя настройка может быть слишком сложной или неоптимальной.
Я обучаю модель классификации изображений (EfficientNet-B3 на изображениях кандзи в оттенках серого) с помощью PyTorch с:
  • AMP (

    Код: Выделить всё

    torch.amp.autocast
    + GradScaler)
  • сбалансированный по классам разделенный набор данных из HDF5
  • Вход: (1, 128, 128) изображений в оттенках серого (~ 620 000 изображений в 3036 классах)
  • Потери: CrossEntropyLoss (сглаживание меток = 0,1)
  • Оптимизатор: AdamW (снижение веса = 1e-4)
Во время обучения я заметил, что моя точность проверки стабилизируется на уровне около 65%.
Для части ArcFace я взял ссылку из:
https://github.com/deepinsight/insightf ... /losses.py
Моя настройка (упрощенная):

Код: Выделить всё

# Source - https://stackoverflow.com/questions/79926183/pytorch-training-issue-with-efficientnetb3-validation-accuracy-plateu
# Posted by Marco
# Retrieved 2026-04-20, License - CC BY-SA 4.0

class ArcFace(nn.Module):
def __init__(self, s=64.0, margin=0.5):
super().__init__()
self.s = s
self.margin = margin
self.cos_m = math.cos(margin)
self.sin_m = math.sin(margin)

def forward(self, logits, labels):
index = torch.where(labels != -1)[0]
target_logit = logits[index, labels[index].view(-1)]

target_logit.arccos_()
logits.arccos_()
final_target_logit = target_logit + self.margin
logits[index, labels[index].view(-1)] = final_target_logit
logits.cos_()

logits = logits * self.s
return logits

class NormLinear(nn.Module):
def __init__(self, in_features, num_classes):
super().__init__()
self.weight = nn.Parameter(torch.FloatTensor(num_classes, in_features))
nn.init.xavier_uniform_(self.weight)

def forward(self, x):
return F.linear(F.normalize(x), F.normalize(self.weight))

Код: Выделить всё

model:
num_classes: 3036
pretrained: True

training:
batch_size: 256
learning_rate: 0.001
epochs: 40
num_workers: 6
shuffle: True
pin_memory: True
persistent_workers: True
prefetch_factor: 4
unfreeze_epoch: 3

optimizer:
type: adamW
weight_decay: 0.0001

data:
...
normalize_mean: [0.5]
normalize_std: [0.5]
val_split: 0.2

device: cuda

...

Код: Выделить всё

class EfficientNetB0Kanji(nn.Module):
def __init__(self, pretrained=True):
super().__init__()
weights = EfficientNet_B0_Weights.IMAGENET1K_V1 if pretrained else None
self.model = efficientnet_b0(weights=weights)

old_weights = None
if pretrained:
old_weights = self.model.features[0][0].weight.clone()

self.model.features[0][0] = nn.Conv2d(1, 32, 3, 2, 1, bias=False)

if pretrained and old_weights is not None:
with torch.no_grad():
self.model.features[0][0].weight[:] = old_weights.mean(dim=1, keepdim=True)

in_features = self.model.classifier[-1].in_features
self.model.classifier[-1] = nn.Linear(in_features, 512)
self.embedding_dim = 512

def forward(self, x):
return self.model(x)

Код: Выделить всё

...
model = EfficientNetB0Kanji(
pretrained=config["model"]["pretrained"],
).to(device)

norm_linear = NormLinear(in_features=512, num_classes=len(dataset.class_to_idx)).to(device)
arcface = ArcFace(s=64.0, margin=0.5)

criterion = nn.CrossEntropyLoss(label_smoothing=0.05)

for param in model.model.features.parameters():
param.requires_grad = False

optimizer = torch.optim.AdamW([
{"params": model.model.classifier.parameters(), "lr": 1e-4},
{"params": norm_linear.parameters(), "lr":  1e-4}
], weight_decay=config["optimizer"]["weight_decay"])

epochs = config["training"]["epochs"]
unfreeze_epoch = config["training"]["unfreeze_epoch"]

scheduler = torch.optim.lr_scheduler.OneCycleLR(
optimizer,
max_lr=[1e-4, 1e-4],
steps_per_epoch=len(train_loader),
epochs=epochs
)

...

Код: Выделить всё

...
for epoch in range(epochs):
if epoch == unfreeze_epoch:
print(f"\nEpoch {epoch + 1}: Unfreezing backbone")
for param in model.model.features.parameters():
param.requires_grad = True

optimizer = torch.optim.AdamW([
{"params": model.model.features.parameters(), "lr": 1e-5},
{"params": model.model.classifier.parameters(), "lr": 1e-4},
{"params": norm_linear.parameters(), "lr": 1e-4}
], weight_decay=config["optimizer"]["weight_decay"])

scheduler = torch.optim.lr_scheduler.OneCycleLR(
optimizer,
max_lr=[1e-5, 1e-4, 1e-4],
steps_per_epoch=len(train_loader),
epochs=epochs - unfreeze_epoch
)

model.train()
...
for step, (images, labels) in loop:
images = images.to(device, non_blocking=True)
labels = labels.to(device, non_blocking=True)

with autocast(device_type='cuda'):
embeddings = model(images)
logits = norm_linear(embeddings)
logits = arcface(logits, labels)
loss = criterion(logits, labels)

scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(
list(model.parameters()) + list(norm_linear.parameters()), max_norm=1.0
)

scaler.step(optimizer)
scaler.update()
scheduler.step()

optimizer.zero_grad(set_to_none=True)
...
with torch.amp.autocast('cuda'):
embeddings = model(images)
logits = norm_linear(embeddings) * arcface.s
val_loss_batch = criterion(logits, labels)

val_loss += val_loss_batch.item()
acc1, acc5 = accuracy(logits, labels, topk=(1, 5))
val_top1_acc += acc1.item()
val_top5_acc += acc5.item()
val_total_batches += 1
...
Что касается моих дополнений, я делаю:

Код: Выделить всё

def get_train_transforms(mean, std):
return v2.Compose([
v2.RandomAffine(
degrees=15,
translate=(0.1, 0.1),
scale=(0.9, 1.1),
shear=5,
fill=0
),

v2.RandomApply(
[v2.ElasticTransform(alpha=20.0, sigma=3.0)],
p=0.2
),

v2.RandomChoice([
v2.GaussianBlur(kernel_size=3, sigma=(0.1, 1.5)),
v2.Identity()
]),

v2.Normalize(mean=mean, std=std)
])

def get_val_transforms(mean, std):
return v2.Compose([
v2.Normalize(mean=mean, std=std)
])
Наблюдения:
ArcFace: Я пытался реализовать это как мог, но по какой-то причине, которую я не могу понять, я сначала начинаю с замороженной фазы с огромными потерями и точностью 0.
Ссылка:
Эпоха 1/30: 100%|██████████| 1898/1898 [04:39

Вернуться в «Python»