# %% [code] {"execution":{"iopub.status.busy":"2021-06-06T09:38:48.898032Z","iopub.execute_input":"2021-06-06T09:38:48.898364Z","iopub.status.idle":"2021-06-06T09:39:17.263214Z","shell.execute_reply.started":"2021-06-06T09:38:48.898333Z","shell.execute_reply":"2021-06-06T09:39:17.260556Z"}}
import numpy as np
import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import os
import torchvision.transforms as transforms

class CustomDatasetImage(Dataset):
    def __init__(self, csv_file, root_dir, transform=None, threechannel=True, image_idx=0, label_idx=1):
        self.customdataDF = pd.read_csv(csv_file) #csv file containing image names　/　画像名が入っているcsvの読み込み
        self.root_dir = root_dir #root directory of where images are stored　/　画像が保存されているdir
        self.transform = transform #transformations to apply to images if any　/　画像の変換
        self.threechannel = threechannel #default is true for standard 3 channel images. false for greyscale images.　/　RGB画像の場合入力が不要、白黒の場合はFalseの入力
        self.image_idx = image_idx
        self.label_idx = label_idx
    
    def __len__(self):
        return len(self.customdataDF) #return the total number of images　/ データセットの数
    
    def __getitem__(self,idx):
        img_path = os.path.join(self.root_dir,self.customdataDF.iloc[idx,self.image_idx]) #obtain the file path of the nth item of the csv / idx個目の画像のファイルパス
        
        if self.threechannel:
            img = Image.open(img_path).convert("RGB") #ensure that the channels are in RGB order / チャンネルはちゃんとRGBにする
        else:
            img = Image.open(img_path).convert("L") #greyscale / 白黒の場合
        
        label = torch.tensor(int(self.customdataDF.iloc[idx,self.label_idx])) #get label and convert to tensor /　ラベルの読み込みとテンソルへの変換
        
        if self.transform: #if there are any transformations, apply　/ 変換があれば実行する
            img=self.transform(img)
        
        return (img,label)

def ObtainMeanStdDataset(csv_file, root_dir, batch_size, device, img_resize=None, num_workers=4, threechannel=True, image_idx=0, label_idx=1):
    #input: csv_file, root_dir, batch_size, device, img_resize, num_workers, threechannel, image_idx, label_idx, batch_size, device　
    # if wish to resize image, pass desired size tuple through img_resize argument
    #　入力：csvファイル名、画像ダイレクトリー、平均・分散を計算するバッチサイズ、デバイス、新画像のサイズ（タプル）、worker数、チャンネル、画像名の行列(csv内)、ラベルの行列(csv内)
    #return mean, std　/　出力：平均、標準偏差
    # VAR(X)= E(X**2)-(E(X))**2
    # std is calculated by taking root of VAR、標準偏差は分散の平方根で計算する
    
    sum_each_mb_channel, sum_each_mb_squared_channel = 0 , 0
    
    if img_resize:
        transform = transforms.Compose([transforms.Resize(img_resize),transforms.ToTensor()]) #if there's a resize, resize
    else:
        transform = transform.ToTensor()
    
    customdata = CustomDatasetImage(csv_file, root_dir, transform, threechannel, image_idx, label_idx)
    loader = DataLoader(customdata, batch_size=batch_size, shuffle=True, num_workers=num_workers)
    
    for minibatch_idx, (img,label) in enumerate(loader):
        img = img.to(device = device)
        sum_each_mb_channel += torch.mean(img, dim=[0,2,3]) # N X C X H x W mean to be taken across N, H and W　/　N,H,Wの平均
        sum_each_mb_squared_channel += torch.mean(img**2, dim=[0,2,3]) #E(X**2)
    
    mean = sum_each_mb_channel/(minibatch_idx+1)#index starts from 0 need to add 1　/　インデックスは０から始まる（ミニバッチ数は+１しないといけない）
    variance = sum_each_mb_squared_channel/(minibatch_idx+1) - mean**2
    std = variance**0.5
    
    return mean, std