Merge pull request #77 from lucasdelimanogueira/tmp

Tmp
This commit is contained in:
lucasdelimanogueira 2024-05-27 21:04:28 -03:00 committed by GitHub
commit c88a8fb8cc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 1 additions and 30 deletions

View file

@ -40,8 +40,6 @@ __host__ void cuda_to_cpu(Tensor* tensor) {
const char* device_str = "cpu";
tensor->device = (char*)malloc(strlen(device_str) + 1);
strcpy(tensor->device, device_str);
printf("Successfully sent tensor to: %s\n", tensor->device);
}
__global__ void add_tensor_cuda_kernel(float* data1, float* data2, float* result_data, int size) {

View file

@ -1,4 +1,4 @@
"""import os
import os
import norch
import norch.distributed as dist
@ -30,30 +30,3 @@ def main():
if __name__ == "__main__":
main()
"""
import norch
import matplotlib.pyplot as plt
import numpy as np
import random
from norch.norchvision import transforms
train_data, test_data = norch.norchvision.datasets.MNIST.splits(transform=transforms.ToTensor())
train_sampler = norch.utils.data.distributed.DistributedSampler(dataset=train_data, num_replicas=10, rank=2)
train_loader = norch.utils.data.Dataloader(train_data, batch_size = 50, sampler=train_sampler)
input_sample, target_sample = train_data[0]
fig = plt.figure(figsize = (20, 10))
columns = 4
rows = 2
# Choose a random image
for image_index, batch in enumerate(train_loader):
image, label = batch
fig.add_subplot(rows, columns, image_index+1)
plt.imshow(np.array(image).reshape(28, 28))
plt.title(label)
plt.axis('off')
if image_index > 6:
break
plt.show()