【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函数

相关推荐
归秋142几秒前
流行音乐编曲软件怎么选:流行音乐创作的实用工具清单
人工智能
Martina_03212 分钟前
同一片森林换个镜头就忽冷忽暖?用6步统一曝光、LUT与相机后处理
人工智能·游戏·数学建模·3d·自然语言处理·aigc·游戏策划
guslegend4 分钟前
多模态 Agent 如何规划 UI 测试路径
人工智能
网络毒刘4 分钟前
MCP 资源与提示(resources/prompts)实战:不只 tools,把只读上下文结构化喂给 Agent
人工智能·ai·cursor
明月_清风5 分钟前
AI Agent 最大的问题,可能不是智商,而是“权限”
人工智能·后端
染指11107 分钟前
135.Agent-多Agent框架-LangChain多智能体(SubAgents子代理)
数据库·人工智能·设计模式·langchain·agent·agents
vx_Biye_Design8 分钟前
expressDeepSeek社团咨询助手的学生社团管理系统60746-计算机课程设计、毕业设计
java·vue.js·spring boot·后端·python·课程设计·express
ACME204421 分钟前
房产商业拍卖适用范围调查:五类物业能否上拍?
人工智能
泡茶喝茶写代码24 分钟前
A股量化数据工程:从 REST 接口到策略信号(第 5 篇):板块成分股映射与对齐
java·python·股票数据api·股票数据·股票数据api接口·股票api数据接口·股票量化数据接口