{"cells":[{"metadata":{},"cell_type":"markdown","source":"## Save RGB Images with 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","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    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","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def 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))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def 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 = dataset + '_rgb'\n    save_file = save_path(dest, experiment, plate, address, site)\n    PIL.Image.fromarray(im_rgb.astype('uint8')).save(save_file)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"os.mkdir(path/'train_rgb')\nfor folder in os.listdir(path/'train'):\n    os.mkdir(path/'train_rgb'/folder)\n    os.mkdir(path/'train_rgb'/folder/'Plate1')\n    os.mkdir(path/'train_rgb'/folder/'Plate2')\n    os.mkdir(path/'train_rgb'/folder/'Plate3')\n    os.mkdir(path/'train_rgb'/folder/'Plate4')\n    \nos.mkdir(path/'test_rgb')\nfor folder in os.listdir(path/'test'):\n    os.mkdir(path/'test_rgb'/folder)\n    os.mkdir(path/'test_rgb'/folder/'Plate1')\n    os.mkdir(path/'test_rgb'/folder/'Plate2')\n    os.mkdir(path/'test_rgb'/folder/'Plate3')\n    os.mkdir(path/'test_rgb'/folder/'Plate4')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"_ = md.apply(lambda row: \n            save_rgb(row['dataset'], row['experiment'], row['plate'], row['well'], row['site']), axis=1)","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}