{"cells":[{"metadata":{},"cell_type":"markdown","source":"## Save RGB Images with Pytorch\n\n**Mostly taken from [Karl Heyer's](https://www.kaggle.com/towardsentropy) script linked [here](https://www.kaggle.com/towardsentropy/save-rgb-image-pytorch)**\n\nThis code saves RGB converted versions of all images in the train and test sets. The conversion is done using code from the [RXRX1 Utils Repo](https://github.com/recursionpharma/rxrx1-utils) adapted to Pytorch.\n\nNote that since this code saves images, it won't work on Kaggle read only kernels"},{"metadata":{"trusted":true},"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport sys\nimport matplotlib.pyplot as plt\nimport torch\nimport PIL\nfrom pathlib import Path\nfrom PIL import Image\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"path = Path('../input')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def pil2tensor(image,dtype):\n    \"Convert PIL style `image` array to torch style image tensor.\"\n    a = np.asarray(image)\n    if a.ndim==2 : a = np.expand_dims(a,2)\n    #a = np.transpose(a, (1, 0, 2))\n    #a = np.transpose(a, (2, 1, 0))\n    a = np.transpose(a, (2, 0, 1))\n    return torch.from_numpy(a.astype(dtype, copy=False) )","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _load_dataset(base_path, dataset, include_controls=True):\n    df =  pd.read_csv(os.path.join(base_path, dataset + '.csv'))\n    if include_controls:\n        controls = pd.read_csv(\n            os.path.join(base_path, dataset + '_controls.csv'))\n        df['well_type'] = 'treatment'\n        df = pd.concat([controls, df], sort=True)\n    df['cell_type'] = df.experiment.str.split(\"-\").apply(lambda a: a[0])\n    df['dataset'] = dataset\n    dfs = []\n    for site in (1, 2):\n        df = df.copy()\n        df['site'] = site\n        dfs.append(df)\n    res = pd.concat(dfs).sort_values(\n        by=['id_code', 'site']).set_index('id_code')\n    return res","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def combine_metadata(base_path=path,\n                     include_controls=True):\n    df = pd.concat(\n        [\n            _load_dataset(\n                base_path, dataset, include_controls=include_controls)\n            for dataset in ['test', 'train']\n        ],\n        sort=True)\n    return df","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"md = combine_metadata()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"md.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def image_path(dataset, experiment, plate,\n               address, site, channel, base_path=path):\n\n    return os.path.join(base_path, dataset, experiment, \"Plate{}\".format(plate),\n                        \"{}_s{}_w{}.png\".format(address, site, channel))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def open_6_channel(dataset, experiment, plate, address, site, base_path=path):\n    return torch.cat([pil2tensor(PIL.Image.open(image_path(dataset, experiment, plate, address, site, i)), np.float32) for i \n            in range(1,7)])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"DEFAULT_CHANNELS = (1, 2, 3, 4, 5, 6)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"RGB_MAP = {\n    1: {\n        'rgb': np.array([19, 0, 249]),\n        'range': [0, 51]\n    },\n    2: {\n        'rgb': np.array([42, 255, 31]),\n        'range': [0, 107]\n    },\n    3: {\n        'rgb': np.array([255, 0, 25]),\n        'range': [0, 64]\n    },\n    4: {\n        'rgb': np.array([45, 255, 252]),\n        'range': [0, 191]\n    },\n    5: {\n        'rgb': np.array([250, 0, 253]),\n        'range': [0, 89]\n    },\n    6: {\n        'rgb': np.array([254, 255, 40]),\n        'range': [0, 191]\n    }\n}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def convert_tensor_to_rgb(t, channels=DEFAULT_CHANNELS, vmax=255, rgb_map=RGB_MAP):\n\n    t = t.permute(1,2,0).numpy()\n    colored_channels = []\n    for i, channel in enumerate(channels):\n        x = (t[:, :, i] / vmax) / \\\n            ((rgb_map[channel]['range'][1] - rgb_map[channel]['range'][0]) / 255) + \\\n            rgb_map[channel]['range'][0] / 255\n        x = np.where(x > 1., 1., x)\n        x_rgb = np.array(\n            np.outer(x, rgb_map[channel]['rgb']).reshape(512, 512, 3),\n            dtype=int)\n        colored_channels.append(x_rgb)\n    im = np.array(np.array(colored_channels).sum(axis=0), dtype=int)\n    im = np.where(im > 255, 255, im)\n    return im\n\ndef save_path(dataset, experiment, plate,\n               address, site, base_path=path):\n\n    return os.path.join(base_path, dataset, experiment, \"Plate{}\".format(plate),\n                        \"{}_s{}.png\".format(address, site))\n\ndef save_rgb(dataset, experiment, plate, address, site):\n    im_6 = open_6_channel(dataset, experiment, plate, address, site)\n    im_rgb = convert_tensor_to_rgb(im_6)\n    dest = 'train_rgb' if 'train' in dataset else 'test_rgb'\n    save_file = save_path(dest, experiment, plate, address, site).replace('../input', '../working/')\n    # print(save_file)\n    PIL.Image.fromarray(im_rgb.astype('uint8')).save(save_file)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"path = ''\nfor type_ in ['train', 'test']:\n    os.mkdir('{}_rgb'.format(type_))\n    for folder in os.listdir('../input/{}'.format(type_)):\n        os.mkdir(os.path.join('{}_rgb'.format(type_), folder))\n        for i in range(1,5):\n            os.mkdir(os.path.join('{}_rgb'.format(type_), folder, 'Plate{}'.format(i)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def plot_experiment(x,y,stater_pos):\n    xn,yn = x,y\n    f,ax = plt.subplots(xn,yn, figsize=(25, 45))\n    for i in range(xn*yn):\n        imgidx = mdtmp.iloc[i+stater_pos] \n        fname = '../working//{}_rgb/{}/Plate{}/{}_s{}.png'.format(imgidx['dataset'], imgidx['experiment'], imgidx['plate'], imgidx['well'], imgidx['site'])\n        a = Image.open(fname) \n        os.remove(fname)\n        ax[int(i/yn), i%yn].imshow(a)\n        ax[int(i/yn), i%yn].title.set_text(imgidx['well_type']+'\\n'+imgidx['well'])\n        ax[int(i/yn), i%yn].title.set_fontsize(10)\n        if imgidx['well_type']=='positive_control':\n            ax[int(i/yn), i%yn].title.set_color('blue')\n        if imgidx['well_type']=='negative_control':\n            ax[int(i/yn), i%yn].title.set_color('red')\n        ax[int(i/yn), i%yn].set_xticks([])\n        ax[int(i/yn), i%yn].set_yticks([])\n    plt.show() ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"EXPERIMENT = 'U2OS-01'\nmdtmp = md[(md['experiment']==EXPERIMENT) & (md['plate']==1) & (md['site']==1)]\n_ = mdtmp.apply(lambda row: save_rgb(row['dataset'], row['experiment'], row['plate'], row['well'], row['site']), axis=1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"imgidx = mdtmp.iloc[0]\nimgidx","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plot_experiment(22,14,0)","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.4","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}