【PyTorch Lightning】

python 复制代码
import os
from torch import optim, nn, utils, Tensor
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor
import lightning as L

# define any number of nn.Modules (or use your current ones)
encoder = nn.Sequential(nn.Linear(28 * 28, 64), nn.ReLU(), nn.Linear(64, 3))
decoder = nn.Sequential(nn.Linear(3, 64), nn.ReLU(), nn.Linear(64, 28 * 28))


# define the LightningModule
class LitAutoEncoder(L.LightningModule):
    def __init__(self, encoder, decoder):
        super().__init__()
        self.encoder = encoder
        self.decoder = decoder

    def training_step(self, batch, batch_idx):
        # training_step defines the train loop.
        # it is independent of forward
        x, _ = batch
        x = x.view(x.size(0), -1)
        z = self.encoder(x)
        x_hat = self.decoder(z)
        loss = nn.functional.mse_loss(x_hat, x)
        # Logging to TensorBoard (if installed) by default
        self.log("train_loss", loss)
        return loss

    def configure_optimizers(self):
        optimizer = optim.Adam(self.parameters(), lr=1e-3)
        return optimizer


# init the autoencoder
autoencoder = LitAutoEncoder(encoder, decoder)

# setup data
dataset = MNIST(os.getcwd(), download=True, transform=ToTensor())
train_loader = utils.data.DataLoader(dataset)

# train the model (hint: here are some helpful Trainer arguments for rapid idea iteration)
trainer = L.Trainer(limit_train_batches=100, max_epochs=1)
trainer.fit(model=autoencoder, train_dataloaders=train_loader)

# load checkpoint
checkpoint = "./lightning_logs/version_0/checkpoints/epoch=0-step=100.ckpt"
autoencoder = LitAutoEncoder.load_from_checkpoint(checkpoint, encoder=encoder, decoder=decoder)

# choose your trained nn.Module
encoder = autoencoder.encoder
encoder.eval()

# embed 4 fake images!
fake_image_batch = torch.rand(4, 28 * 28, device=autoencoder.device)
embeddings = encoder(fake_image_batch)
print("⚡" * 20, "\nPredictions (4 image embeddings):\n", embeddings, "\n", "⚡" * 20)
python 复制代码
tensorboard --logdir .

Module中要定义forward函数;

Lightning Module中除了forward函数,还要定义configure_optimizers, training_step和validation_step函数

相关推荐
iori97king几秒前
织信开发日志 18:从 informat-skills 看织信如何把平台能力交给 AI Agent
人工智能·低代码·织信
daxiang_ipo几秒前
AI应用,正式步入质变跃迁期
人工智能·搜索引擎·百度
前沿在线1 分钟前
2026世界机器人大会主论坛大咖观点(一)
人工智能·ai·大模型
智搜广告2 分钟前
AI回答优化公司智搜广告让品牌成为推荐首选
大数据·人工智能·python·elasticsearch·geo
Three_ST3 分钟前
沐神-动手学习深度学习-习题 4.3. 多层感知机的简洁实现
人工智能·pytorch·python·深度学习·机器学习
必须会一定会3 分钟前
Agent Handoff M5 发布验收:`CHANGES.md`、`npm pack`、`release:check` 与干净环境安装验证
前端·人工智能·npm·node.js·ai编程
Allen_LVyingbo5 分钟前
面向电子病历的批量语义分析自动化工具:从设计到实战(上)
运维·人工智能·机器学习·语言模型·自然语言处理·自动化·健康医疗
明志数科6 分钟前
从实验室到工厂产线:具身智能训练数据的环境差异、分布偏移与工业级采集方案
人工智能·深度学习·计算机视觉
武子康8 分钟前
我让 Qwen3.6-27B 真改了一次 Git 仓库:工具调用怎样形成 Agent 闭环
人工智能·后端·agent
嘻嘻的AI日记8 分钟前
告别会后整理负担|智能会议系统,实现会纪要自动生成
人工智能