Обучение PyTorch SRGAN чрезвычайно медленно на процессоре для сверхвысокого разрешения изображения – как оптимизировать?Python

Программы на Python
Anonymous
Обучение PyTorch SRGAN чрезвычайно медленно на процессоре для сверхвысокого разрешения изображения – как оптимизировать?

Сообщение Anonymous »

  • Я работаю над задачей суперразрешения листьев растений, используя PyTorch. Я построил модель на основе SRGAN с сетями генератора и дискриминатора для сверхвысокого разрешения изображений.
  • Модель берет изображения листьев растений с низким разрешением и генерирует выходные данные с высоким разрешением. Во время обучения процесс занимает слишком много времени, потому что я запускаю его на процессоре, а не на графическом процессоре.
  • Каждая эпоха выполняется очень медленно из-за вычислений генератора, дискриминатора и потерь восприятия. Я использую методы смешанной точности и оптимизации, но производительность ЦП по-прежнему низкая.
  • Я хочу знать, как сократить время обучения и оптимизировать этот код для выполнения ЦП. Будут полезны предложения по повышению скорости, сокращению эпох или упрощению модели.
Код:
best_loss = float('inf')

for epoch in range(1, NUM_EPOCHS + 1):
G.train(); D.train()
g_losses, d_losses = [], []

for lr_img, hr_img in train_dl:
lr_img = lr_img.to(DEVICE)
hr_img = hr_img.to(DEVICE)
B = lr_img.size(0)

real_label = torch.ones (B, 1, 1, 1, device=DEVICE) * 0.9 # label smoothing
fake_label = torch.zeros(B, 1, 1, 1, device=DEVICE) + 0.1

# -- Discriminator --
opt_D.zero_grad()
with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):
fake_hr = G(lr_img).detach()
d_real = D(hr_img)
d_fake = D(fake_hr)
# Match spatial size of labels to discriminator output
rl = real_label.expand_as(d_real)
fl = fake_label.expand_as(d_fake)
loss_D = (criterion_adv(d_real, rl) + criterion_adv(d_fake, fl)) * 0.5

scaler.scale(loss_D).backward()
scaler.step(opt_D)

# -- Generator --
opt_G.zero_grad()
with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):
fake_hr = G(lr_img)
d_fake = D(fake_hr)
rl = real_label.expand_as(d_fake)

loss_pix = criterion_pix(fake_hr, hr_img)
loss_adv = criterion_adv(d_fake, rl)

if perc_net is not None:
with torch.no_grad():
feat_real = perc_net(hr_img)
feat_fake = perc_net(fake_hr)
loss_perc = F.l1_loss(feat_fake, feat_real.detach())
else:
loss_perc = torch.tensor(0.0, device=DEVICE)

loss_G = (LAMBDA_PIX * loss_pix +
LAMBDA_ADV * loss_adv +
LAMBDA_PERC * loss_perc)

scaler.scale(loss_G).backward()
scaler.step(opt_G)
scaler.update()

g_losses.append(loss_G.item())
d_losses.append(loss_D.item())

sched_G.step(); sched_D.step()

mean_g = np.mean(g_losses)
mean_d = np.mean(d_losses)

if mean_g < best_loss:
best_loss = mean_g
torch.save(G.state_dict(), 'best_generator.pth')

if epoch % 10 == 0 or epoch == 1:
print(f'Epoch [{epoch:>3}/{NUM_EPOCHS}] G: {mean_g:.4f} D: {mean_d:.4f}')

print('Training complete. Best G loss:', best_loss)

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