PyNorch/examples/train_multigpu.py

107 lines
2.8 KiB
Python
Raw Permalink Normal View History

2024-06-03 15:37:57 -03:00
import os
import norch
import norch.distributed as dist
import norch.distributed
import norch.nn as nn
import norch.optim as optim
from norch.nn.parallel import DistributedDataParallel
from norch.utils.data.distributed import DistributedSampler
from norch.norchvision import transforms as T
import random
random.seed(1)
def main():
local_rank = int(os.getenv('OMPI_COMM_WORLD_LOCAL_RANK', -1))
rank = int(os.getenv('OMPI_COMM_WORLD_RANK', -1))
world_size = int(os.getenv('OMPI_COMM_WORLD_SIZE', -1))
dist.init_process_group(
rank,
world_size
)
BATCH_SIZE = 32
device = local_rank
epochs = 10
transform = T.Compose(
[
T.ToTensor(),
T.Reshape([-1, 784, 1])
]
)
target_transform = T.Compose(
[
T.ToTensor()
]
)
2024-06-05 15:43:19 -03:00
print("Loading data")
2024-06-03 15:37:57 -03:00
train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transform, target_transform=target_transform)
distributed_sampler = DistributedSampler(dataset=train_data, num_replicas=world_size, rank=rank)
train_loader = norch.utils.data.DataLoader(train_data, batch_size=BATCH_SIZE, sampler=distributed_sampler)
class MyModel(nn.Module):
def __init__(self):
super(MyModel, self).__init__()
self.fc1 = nn.Linear(784, 30)
self.sigmoid1 = nn.Sigmoid()
self.fc2 = nn.Linear(30, 10)
self.sigmoid2 = nn.Sigmoid()
def forward(self, x):
out = self.fc1(x)
out = self.sigmoid1(out)
out = self.fc2(out)
out = self.sigmoid2(out)
return out
2024-06-05 15:43:19 -03:00
print("Creating model")
2024-06-03 15:37:57 -03:00
model = MyModel().to(device)
model = DistributedDataParallel(model)
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.01)
loss_list = []
print(f"Starting training on Rank {rank}/{world_size}\n\n")
for epoch in range(epochs):
2024-06-05 15:43:19 -03:00
avg_loss = 0
num_steps = 0
2024-06-03 15:37:57 -03:00
for idx, batch in enumerate(train_loader):
2024-06-05 15:43:19 -03:00
if idx % 100 == 0 and rank == 0:
print(f"Epoch: {epoch}/{epochs} - Step: {idx} / {len(train_loader)}")
2024-06-03 15:37:57 -03:00
inputs, target = batch
inputs = inputs.to(device)
target = target.to(device)
outputs = model(inputs)
loss = criterion(outputs, target)
optimizer.zero_grad()
loss.backward()
optimizer.step()
2024-06-05 15:43:19 -03:00
avg_loss += loss[0]
num_steps += 1
avg_loss = avg_loss / num_steps
2024-06-03 15:37:57 -03:00
if rank == 0:
2024-06-05 15:43:19 -03:00
print(f'Epoch [{epoch + 1}/{epochs}], Loss: {avg_loss:.4f}')
loss_list.append(avg_loss)
2024-06-03 15:37:57 -03:00
if __name__ == "__main__":
main()