[BUG]: 使用 use_fp8=true 的 LowLevelZeroPlugin 会导致 `形状不匹配错误: 形状 [32, 512] 对于大小为 512 的输入是无效的`
class RandomDataset(Dataset): def init(self, num_samples=32 * 100, input_dim=1024, num_classes=10): self.x = torch.randn(num_samples, input_dim) self.y = torch.randint(0, num_classes, (num_samples,))
def len(self): return len(self.x)
def getitem(self, idx): return self.x[idx], self.y[idx]
class MLP(nn.Module): def init(self, input_dim=1024, hidden_dim=512, num_layers=10, num_classes=10): super().init() layers = [] for i in range(num_layers): in_dim = input_dim if i == 0 else hidden_dim layers.append(nn.Linear(in_dim, hidden_dim)) layers.append(nn.ReLU()) layers.append(nn.Linear(hidden_dim, num_classes)) self.net = nn.Sequential(*layers)
def forward(self, x): return self.net(x)
def main(): seed = 1024 colossalai.launch_from_torch(seed=seed) plugin = LowLevelZeroPlugin( use_fp8=True )
booster = Booster(plugin=plugin)
model = MLP() optimizer = optim.Adam(model.parameters(), lr=1e-3) criterion = nn.CrossEntropyLoss()
dataset = RandomDataset() train_dataloader = DataLoader(dataset, batch_size=32, shuffle=False)
model, optimizer, criterion, train_dataloader, _ = booster.boost(model, optimizer, criterion, train_dataloader)
precision = getattr(plugin, "precision", "fp16") dtype_map = {"fp16": torch.float16, "bf16": torch.bfloat16, "fp32": torch.float32} dtype = dtype_map.get(precision, torch.float16)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.train() for epoch in range(1): total_loss = 0 for step, (x, y) in enumerate(train_dataloader):
x = x.to(device=device, dtype=dtype) y = y.to(device=device)
optimizer.zero_grad() output = model(x) loss = criterion(output, y)
booster.backward(loss, optimizer) optimizer.step()
total_loss += loss.item()
内容来源: hpcaitech/ColossalAI