MCPcopy Create free account
hub / github.com/pytorch/executorch / get_dataset

Function get_dataset

examples/samsung/scripts/vit.py:27–62  ·  view source on GitHub ↗
(dataset_path, data_size)

Source from the content-addressed store, hash-verified

25
26
27def get_dataset(dataset_path, data_size):
28 from torchvision import datasets, transforms
29
30 image_shape = (256, 256)
31 crop_size = 224
32 shuffle = True
33
34 def get_data_loader():
35 preprocess = transforms.Compose(
36 [
37 transforms.Resize(image_shape),
38 transforms.CenterCrop(crop_size),
39 transforms.ToTensor(),
40 transforms.Normalize(
41 mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]
42 ),
43 ]
44 )
45 imagenet_data = datasets.ImageFolder(dataset_path, transform=preprocess)
46 return torch.utils.data.DataLoader(
47 imagenet_data,
48 shuffle=shuffle,
49 )
50
51 # prepare input data
52 inputs, targets, input_list = [], [], ""
53 data_loader = get_data_loader()
54 for index, data in enumerate(data_loader):
55 if index >= data_size:
56 break
57 feature, target = data
58 inputs.append((feature,))
59 targets.append(target)
60 input_list += f"input_{index}_0.bin\n"
61
62 return inputs, targets, input_list
63
64
65if __name__ == "__main__":

Callers 1

vit.pyFile · 0.70

Calls 2

get_data_loaderFunction · 0.70
appendMethod · 0.45

Tested by

no test coverage detected