# pl.LightningModule, pl.LightningDataModule, pl.Trainer
import os
import torch
from torch import utils
from torch import optim, nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
import torchvision
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor
import pytorch_lightning as pl
from pytorch_lightning.plugins import DDPPlugin
from pytorch_lightning.loggers import WandbLogger
#### WandbLogger ####
wandb_logger = WandbLogger(project="mnist-vae", name="mnist-vae-v1")
#### pl.LightningDataModule ####
class MNISTDataModule(pl.LightningDataModule):
def __init__(self, data_dir="./data", batch_size=32, num_workers=8):
super().__init__()
self.data_dir = data_dir
self.batch_size = batch_size
self.num_workers = num_workers
self.transform = ToTensor()
def prepare_data(self):
# download, split, etc...
# only called on 1 GPU/TPU in distributed
# 下载数据集
MNIST(self.data_dir, train=True, download=True)
MNIST(self.data_dir, train=False, download=True)
def setup(self, stage=None):
# make assignments here (val/train/test split)
# called on every process in DDP
# stage = None, fit, validate, test, predict
# 分割数据集为训练集和验证集
if stage == "fit" or stage is None:
self.mnist_train = MNIST(
self.data_dir, train=True, transform=self.transform
)
self.mnist_val = MNIST(self.data_dir, train=False, transform=self.transform)
# 如果有测试阶段,可以在这里处理测试数据
if stage == "test" or stage is None:
self.mnist_test = MNIST(
self.data_dir, train=False, transform=self.transform
)
def train_dataloader(self):
return DataLoader(
self.mnist_train,
batch_size=self.batch_size,
shuffle=True,
num_workers=self.num_workers,
)
def val_dataloader(self):
return DataLoader(
self.mnist_val,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
)
def test_dataloader(self):
return DataLoader(
self.mnist_test,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
)
data_module = MNISTDataModule()
print("datamodule load done!")
#### pl.LightningModule ####
"""
forward: define a forward pass(input batch data, output batch data)
train_step: define a training step(input batch data, output loss)
val_step: define a validation step(input batch data, output loss)
test_step: define a test step(input batch data, output loss)
configure_optimizers: define optimizer and lr_scheduler
"""
# define the LightningModule
class LitAutoEncoder(pl.LightningModule):
def __init__(self):
super().__init__()
self.encoder = nn.Sequential(
nn.Linear(28 * 28, 64), nn.ReLU(), nn.Linear(64, 3)
)
self.decoder = nn.Sequential(
nn.Linear(3, 64), nn.ReLU(), nn.Linear(64, 28 * 28)
)
def forward(self, x):
x = x.view(x.size(0), -1) # 确保展平操作正确
z = self.encoder(x)
x_hat = self.decoder(z)
return x_hat
def common_step(self, batch, batch_idx, optimizer_idx=None):
x, y = batch
x_hat = self.forward(x)
loss = F.mse_loss(
x_hat, x.view(x.size(0), -1)
) # 确保损失计算中输入和输出尺寸匹配
return loss
def training_step(self, batch, batch_idx):
loss, x_hat = self.common_step(batch, batch_idx)
self.log("train_step_loss", loss)
if batch_idx % 100 == 0:
# 生成图像网格
grid = make_grid(x_hat.view(-1, 1, 28, 28), nrow=8)
# 使用WandB记录图像
wandb_logger.log({"train_images": [wandb.Image(grid, caption="Train Images")], "train_step_loss": loss.item()}, step=self.global_step)
return {"loss": loss}
def configure_optimizers(self):
optimizer = optim.Adam(self.parameters(), lr=1e-3)
return optimizer
def validation_step(self, batch, batch_idx):
loss = self.common_step(batch, batch_idx)
self.log("val_step_loss", loss)
return loss
def predict_step(self, batch, batch_idx):
x, y = batch
x_hat = self.forward(x)
pred = argmax(x_hat, dim=1)
return pred
# init the autoencoder
autoencoder = LitAutoEncoder()
print("model create done!")
#### pl.Trainer ####
# 创建DDPPlugin实例,并设置find_unused_parameters=True
ddp_plugin = DDPPlugin(find_unused_parameters=True)
# train the model with DDP strategy and 2 GPUs
trainer = pl.Trainer(
strategy=ddp_plugin, # DDPPlugin or other training type plugins.
accelerator="gpu", # accelerator types ("cpu", "gpu", "tpu", "ipu", "auto")
devices=2, # gpu nums
# accumulate_grad_batches=4, # Accumulates grads every k batches
logger=wandb_logger,
limit_train_batches=100,
max_epochs=5,
)
print("trainer create done!")
trainer.fit(model=autoencoder, datamodule=data_module)
pytorch lightning base demo (pl.LightningModule, pl.LightningDataModule, pl.Trainer)
最新推荐文章于 2026-07-03 09:26:50 发布
该文章已生成可运行项目,
本文章已经生成可运行项目

4220

被折叠的 条评论
为什么被折叠?



