mindspore/raf.py

52 lines
1.7 KiB
Python

import cv2
import numpy as np
from PIL import Image
import os
import copy
import torchvision
import torch
from .randaugment import RandAugment
class TransformTwice:
def __init__(self, transform):
self.transform = transform
self.strong_transfrom = copy.deepcopy(transform)
self.strong_transfrom.transforms.insert(0, RandAugment(3,5))
def __call__(self, inp):
out1 = self.transform(inp)
out2 = self.transform(inp)
out3 = self.strong_transfrom(inp)
out 4= self.strong_transfrom(inp)
return out1, out2, out3
def get_raf(train_root, train_file_list, test_root, test_file_list, n_labeled, transform_train=None, transform_val=None):
train_labeled_idxs, train_unlabeled_idxs = data_split(train_file_list, int(n_labeled))
train_labeled_dataset = Dataset_RAF_labeled(train_root, train_file_list, train_labeled_idxs, transform=transform_train)
train_unlabeled_dataset = Dataset_RAF_unlabeled(train_root, train_file_list, train_unlabeled_idxs, transform=TransformTwice(transform_train))
test_dataset = Dataset_RAF(test_root, test_file_list, transform=transform_val)
print (f"#Labeled: {len(train_labeled_idxs)} #Unlabeled: {len(train_unlabeled_dataset)}")
return train_labeled_dataset, train_unlabeled_dataset, test_dataset
def target_read(path):
label_list = []
with open(path) as f:
img_label_list = f.read().splitlines()
for info in img_label_list:
_, label_name = info.split(' ')
label_list.append(int(label_name))
return label_list
def data_split(filename, n_labeled):
labels = target_read(filename)
labels = np.array(labels)
train_labeled_idxs = []
train_unlabeled_idxs = []