{"cells":[{"metadata":{"_uuid":"2cba857b734cb99ef2270825c46bde046c0172de"},"cell_type":"markdown","source":"Hi Kaggle!\n\nIn this notebook I share with you the tool I use to work with PLAsTiCC data. \n\nBelow, you'll find the definition of the module, which you can move to another file.\n\nIn the next cells, I'll show you what's the module can show you :) Enjoy!"},{"metadata":{"trusted":true,"_uuid":"62cc14a0b1fe63928204eee51543390feee79077","_kg_hide-input":true,"_kg_hide-output":false},"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\nimport warnings\n\nwarnings.filterwarnings('ignore')\nimport itertools\n\n\nclass PlasticcVis:\n    def __init__(self, data):\n        self.data = data\n\n        # Fix for missing PROJ4 env var https://github.com/conda-forge/basemap-feedstock/issues/30#issuecomment-423512069\n        import os\n        import conda\n\n        conda_file_dir = conda.__file__\n        conda_dir = conda_file_dir.split('lib')[0]\n        proj_lib = os.path.join(os.path.join(conda_dir, 'share'), 'proj')\n        os.environ[\"PROJ_LIB\"] = proj_lib\n\n    def light_curve(self, object_id):\n        \"\"\"\n        Plot light curves of a single object in all passbands\n        :param object_id: id of the object for which the curves will be plotted\n        :return: ax with drown chart\n        \"\"\"\n        ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"passband\", data=self.data[self.data['object_id'] == object_id])\n        ax.plot()\n        return ax\n\n    def light_curve_split(self, object_id):\n        \"\"\"\n        Plot light curves of a single object with separate chart for each passband\n        :param object_id: id of the object for which the curves will be plotted\n        :return: axes array with plotted charts\n        \"\"\"\n        data = self.data[self.data['object_id'] == object_id]\n        passbands = data['passband'].sort_values().unique()\n        fig, axes = plt.subplots(nrows=len(passbands), ncols=1, sharex=True, sharey=True)\n        fig.tight_layout()\n        with sns.plotting_context(\"notebook\", rc={\"lines.linewidth\": 3}):\n            for i, passband in enumerate(passbands):\n                ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"object_id\", data=data[data['passband'] == passband],\n                                  ax=axes[i], legend=False)\n                ax.set_title(f\"Passband {passband}\")\n\n        return axes\n\n    def light_curve_compared(self, object_id, target):\n        \"\"\"\n        Plot specified object light curves along with other light curves in specified target class\n        :param object_id: id of the object for which the curves will be plotted\n        :param target: target class of other objects which will be drawn along specified object\n        :return: ax with drown chart\n        \"\"\"\n        ax = None\n        other_ids = self.data[(self.data['object_id'] != object_id) & (self.data['target'] == target)][\n            'object_id'].unique()\n        for other_id in other_ids:\n            ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"passband\", data=self.data[self.data['object_id'] == other_id],\n                              palette=sns.cubehelix_palette(dark=.7, light=.9, as_cmap=True),\n                              legend='brief' if ax is None else False)\n        with sns.plotting_context(\"notebook\", rc={\"lines.linewidth\": 3}):\n            ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"passband\", data=self.data[self.data['object_id'] == object_id],\n                              palette=sns.light_palette(\"green\", as_cmap=True))\n        ax.plot()\n\n        return ax\n\n    def light_curve_split_compared(self, object_id, target):\n        \"\"\"\n        Plot light curves of a single object with separate chart for each passband\n        along with other light curves in specified target class\n        :param object_id: id of the object for which the curves will be plotted\n        :param target: target class of other objects which will be drawn along specified object\n        :return: axes array with plotted charts\n        \"\"\"\n        passbands = self.data['passband'].sort_values().unique()\n        fig, axes = plt.subplots(nrows=len(passbands), ncols=1, sharex=True, sharey=True)\n        fig.tight_layout()\n        for i, passband in enumerate(passbands):\n            ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"object_id\", data=self.data[\n                (self.data['object_id'] != object_id) & (self.data['target'] == target) & (\n                            self.data['passband'] == passband)], ax=axes[i])\n            with sns.plotting_context(\"notebook\", rc={\"lines.linewidth\": 3}):\n                sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"object_id\", data=self.data[\n                    (self.data['object_id'] == object_id) & (self.data['passband'] == passband)],\n                             palette=sns.light_palette(\"green\", reverse=True, as_cmap=True), ax=ax)\n            ax.set_title(f\"Passband {passband}\")\n            ax.plot()\n\n        return axes\n\n    def light_curve_interpolate_compare(self, object_id, target):\n        \"\"\"\n        Plot light curves of a single object along\n        with interpolated line representing average light curve of the target class\n        :param object_id: id of the object for which the curves will be plotted\n        :param target: target class of other objects which will be drawn along specified object\n        :return: ax with drown chart\n        \"\"\"\n        passbands = self.data['passband'].sort_values().unique()\n        fig, ax = plt.subplots()\n        for i, passband in enumerate(passbands):\n            data = self.data[(self.data['passband'] == passband) & (self.data['target'] == target) & (\n                        self.data['object_id'] != object_id)]\n            x_pane = np.linspace(min(data['mjd']), max(data['mjd']), 100)\n            y_pane = np.poly1d(np.polyfit(data['mjd'], data['flux'], 15, 25))(x_pane)\n            sns.lineplot(x=x_pane, y=y_pane, palette=sns.dark_palette(\"palegreen\", as_cmap=True), ax=ax)\n\n            with sns.plotting_context(\"notebook\", rc={\"lines.linewidth\": 3}):\n                sns.lineplot(x=\"mjd\", y=\"flux\", data=self.data[\n                    (self.data['object_id'] == object_id) & (self.data['passband'] == passband)], ax=ax,\n                             palette=sns.light_palette(\"green\", reverse=True, as_cmap=True))\n            ax.plot()\n\n        return ax\n\n    def light_curve_interpolate_split_compare(self, object_id, target):\n        \"\"\"\n        Plot light curves of a single object with separate chart for each passband\n        along with interpolated line representing average3 light curve of the target class\n        :param object_id: id of the object for which the curves will be plotted\n        :param target: target class of other objects which interpolation will be calculated\n        :return: axes array with plotted charts\n        \"\"\"\n        passbands = self.data['passband'].sort_values().unique()\n        fig, axes = plt.subplots(nrows=len(passbands), ncols=1, sharex=True, sharey=False)\n        for i, passband in enumerate(passbands):\n            ax = axes[i]\n            data = self.data[(self.data['passband'] == passband) & (self.data['target'] == target) & (\n                        self.data['object_id'] != object_id)]\n            x_pane = np.linspace(min(data['mjd']), max(data['mjd']), 100)\n            y_pane = np.poly1d(np.polyfit(data['mjd'], data['flux'], 15, 25))(x_pane)\n            sns.lineplot(x=x_pane, y=y_pane, palette=sns.dark_palette(\"palegreen\", as_cmap=True), ax=ax)\n\n            with sns.plotting_context(\"notebook\", rc={\"lines.linewidth\": 3}):\n                sns.lineplot(x=\"mjd\", y=\"flux\", data=self.data[\n                    (self.data['object_id'] == object_id) & (self.data['passband'] == passband)], ax=ax,\n                             palette=sns.dark_palette(\"green\", reverse=True, as_cmap=True))\n            ax.set_title(f\"Passband {passband}\")\n            ax.plot()\n\n        return axes\n\n    def sky_pos(self, object_id):\n        \"\"\"\n        Show position of specified object on Aitoff projection of night sky between other objects from dataset\n        :param object_id: id of the object for which the curves will be plotted\n        :return: ax with drown chart\n        \"\"\"\n        from mpl_toolkits.basemap import Basemap\n\n        fig, ax = plt.subplots()\n        data = self.data.copy()\n        m = Basemap(projection='hammer', lon_0=0, ax=ax)\n        data['x'], data['y'] = m(data['ra'].apply(lambda x: x - 180).tolist(), data['decl'].tolist())\n\n        sns.scatterplot(x='x', y='y', hue=\"object_id\", data=data[data['object_id'] != object_id],\n                        palette=sns.cubehelix_palette(dark=.6, light=.9, as_cmap=True), ax=ax, linewidth=0)\n        sns.scatterplot(x=\"x\", y=\"y\", hue=\"object_id\", data=data[data['object_id'] == object_id], ax=ax,\n                        palette=sns.light_palette(\"palegreen\", as_cmap=True), marker='x', linewidth=4, s=400)\n\n        m.imshow(plt.imread('../input/plasticc-merged/mell_rgb_450.jpg', 0))\n        ax.set_title(f'Sky placement of the object {object_id} among other objects')\n        ax.legend(loc='center right', bbox_to_anchor=(1.1, 0.9), ncol=1)\n        fig.show()\n\n        return ax\n\n    def sky_pos_target(self, object_id, target):\n        \"\"\"\n        Show position of specified object on Aitoff projection of night sky between other objects in specified target\n        :param object_id: id of the object for which the curves will be plotted\n        :return: ax with drown chart\n        \"\"\"\n        from mpl_toolkits.basemap import Basemap\n        fig, ax = plt.subplots()\n        data = self.data.copy()\n        m = Basemap(projection='hammer', lon_0=0, ax=ax)\n        data['x'], data['y'] = m(data['ra'].apply(lambda x: x - 180).tolist(), data['decl'].tolist())\n\n        sns.scatterplot(x='x', y='y', hue=\"object_id\",\n                        data=data[(data['object_id'] != object_id) & (self.data['target'] == target)],\n                        palette=sns.cubehelix_palette(dark=.6, light=.9, as_cmap=True), ax=ax, linewidth=0)\n        sns.scatterplot(x=\"x\", y=\"y\", hue=\"object_id\", data=data[data['object_id'] == object_id], ax=ax,\n                        palette=sns.light_palette(\"palegreen\", as_cmap=True), marker='x', linewidth=4, s=400)\n\n        m.imshow(plt.imread('../input/plasticc-merged/mell_rgb_450.jpg', 0))\n        ax.set_title(f'Sky placement of the object {object_id} among other objects in target {target}')\n        ax.legend(loc='center right', bbox_to_anchor=(1.1, 0.9), ncol=1)\n        fig.show()\n        return ax\n\n    def galactic_pos(self, object_id):\n        \"\"\"\n        Show position of specified object on Aitoff projection of our galaxy between other objects from dataset\n        :param object_id: id of the object for which the curves will be plotted\n        :return: ax with drown chart\n        \"\"\"\n        from mpl_toolkits.basemap import Basemap\n        fig, ax = plt.subplots()\n        data = self.data.copy()\n        m = Basemap(projection='hammer', lon_0=0, ax=ax)\n        data['x'], data['y'] = m(data['gal_l'].apply(lambda x: x - 180).tolist(), data['gal_b'].tolist())\n\n        sns.scatterplot(x='x', y='y', hue=\"object_id\", data=data[data['object_id'] != object_id],\n                        palette=sns.cubehelix_palette(dark=.6, light=.9, as_cmap=True), ax=ax, linewidth=0)\n        sns.scatterplot(x=\"x\", y=\"y\", hue=\"object_id\", data=data[data['object_id'] == object_id], ax=ax,\n                        palette=sns.light_palette(\"palegreen\", as_cmap=True), marker='x', linewidth=4, s=400)\n\n        m.imshow(plt.imread('../input/plasticc-merged/Milky_Way_infrared.jpg', 0))\n        ax.set_title(f'Galactic placement of the object {object_id} among other objects')\n        ax.legend(loc='center right', bbox_to_anchor=(1.1, 0.9), ncol=1)\n        fig.show()\n        return ax\n\n    def galactic_pos_target(self, object_id, target):\n        \"\"\"\n        Show position of specified object on Aitoff projection of our galaxy between other objects in specified target\n        :param object_id: id of the object for which the curves will be plotted\n        :return: ax with drown chart\n        \"\"\"\n        from mpl_toolkits.basemap import Basemap\n        fig, ax = plt.subplots()\n        data = self.data.copy()\n        m = Basemap(projection='hammer', lon_0=0, ax=ax)\n        data['x'], data['y'] = m(data['gal_l'].apply(lambda x: x - 180).tolist(), data['gal_b'].tolist())\n\n        sns.scatterplot(x='x', y='y', hue=\"object_id\",\n                        data=data[(data['object_id'] != object_id) & (self.data['target'] == target)],\n                        palette=sns.cubehelix_palette(dark=.6, light=.9, as_cmap=True), ax=ax, linewidth=0)\n        sns.scatterplot(x=\"x\", y=\"y\", hue=\"object_id\", data=data[data['object_id'] == object_id], ax=ax,\n                        palette=sns.light_palette(\"palegreen\", as_cmap=True), marker='x', linewidth=4, s=400)\n\n        m.imshow(plt.imread('../input/plasticc-merged/Milky_Way_infrared.jpg', 0))\n        ax.set_title(f'Galactic placement of the object {object_id} among other objects in target {target}')\n        ax.legend(loc='center right', bbox_to_anchor=(1.1, 0.9), ncol=1)\n        fig.show()\n        return ax\n\n    def passband_target_interpolate_matrix(self):\n        \"\"\"\n        Calculates interpolated line for every target and every passband, then plots them in separate charts\n        :return: axes array with plotted charts\n        \"\"\"\n        targets = self.data['target'].sort_values().unique()\n        passbands = self.data['passband'].sort_values().unique()\n        fig, axes = plt.subplots(nrows=len(targets), ncols=len(passbands), sharex=True, sharey=False)\n        axes = axes.flatten()\n        fig.tight_layout()\n        for i, (target, passband) in enumerate(list(itertools.product(targets, passbands))):\n            data = self.data[(self.data['passband'] == passband) & (self.data['target'] == target)]\n            x_pane = np.linspace(min(data['mjd']), max(data['mjd']), 100)\n            y_pane = np.poly1d(np.polyfit(data['mjd'], data['flux'], 15, 25))(x_pane)\n            ax = sns.lineplot(x=x_pane, y=y_pane, palette=sns.dark_palette(\"palegreen\", as_cmap=True), ax=axes[i])\n            ax.set_title(f\"Passband {passband}, target {target}\")\n\n        return axes\n\n    def extragalactic_light_curves_split(self):\n        \"\"\"\n        Plot every light curve from extragalactic sources, split into targets and passbands\n        :return: axes array with plotted charts\n        \"\"\"\n        sources = [15, 42, 52, 62, 64, 67, 88, 90, 95]\n        passbands = self.data['passband'].sort_values().unique()\n        pairs = list(itertools.product(sources, passbands))\n        fig, axes = plt.subplots(nrows=len(sources), ncols=len(passbands), sharex=True, sharey=False)\n        axes = axes.flatten()\n        for i, (source, passband) in enumerate(pairs):\n            ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"object_id\",\n                              data=self.data[(self.data['passband'] == passband) & (self.data['target'] == source)],\n                              ax=axes[i], legend=False)\n            ax.set_title(f\"source {source} passband {passband}\")\n\n    def galactic_light_curves_split(self):\n        \"\"\"\n        Plot every light curve from galactic sources, split into targets and passbands\n        :return: axes array with plotted charts\n        \"\"\"\n        sources = [6, 16, 53, 65, 92]\n        passbands = self.data['passband'].sort_values().unique()\n        pairs = list(itertools.product(sources, passbands))\n        fig, axes = plt.subplots(nrows=len(sources), ncols=len(passbands), sharex=True, sharey=False)\n        axes = axes.flatten()\n        for i, (source, passband) in enumerate(pairs):\n            ax = sns.lineplot(x=\"mjd\", y=\"flux\", hue=\"object_id\",\n                              data=self.data[(self.data['passband'] == passband) & (self.data['target'] == source)],\n                              ax=axes[i], legend=False)\n            ax.set_title(f\"source {source} passband {passband}\")","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"5649f649cce0b0076b2d1aec8f98378acc26f5dc"},"cell_type":"markdown","source":"# How to use PlasticcVis module\n\nLet's assume we have object of interest with id 80067866 and it's target class is 6.\nWe want to find out more about it. \nThis module makes this analysis easy.\n\n\n\nFirst, let's initialize our module with data. Remember that this has to be combined dataset, containing both flux and meta data"},{"metadata":{"trusted":true,"_uuid":"4e35ea1a368622fea0dc404f6ba55fcb8e0addb0"},"cell_type":"code","source":"plt.rcParams[\"figure.figsize\"] = (20, 10)\nimport pandas as pd\n\nanalysis = PlasticcVis(data=pd.read_csv('../input/plasticc-merged/dataset.csv'))\nanalysis.data.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"37a96b71f1ecc70e7f0d1a90653134fb1b57c455"},"cell_type":"code","source":"object_id = 80067866\ntarget = 6","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"5c7d5820a98ee5974c97cff0cd8027b385744df3"},"cell_type":"markdown","source":"Now, let's plot the light curve of our object"},{"metadata":{"trusted":true,"_uuid":"f79e41d0b2065e97d4c8c2603cc60a7e052382a6"},"cell_type":"code","source":"ax = analysis.light_curve(object_id=object_id)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4435054f5241282188ca085e3beeb1ef6189d1d8"},"cell_type":"markdown","source":"Now, how that compares to other objects in it's target, is it an average object or some outlier?"},{"metadata":{"trusted":true,"_uuid":"26d94664aedfe3d389fe20ba177bc50ad7a1b71b"},"cell_type":"code","source":"ax = analysis.light_curve_compared(object_id=object_id, target=target)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"cd0403a06b3b1a51fbf1226b4ff8680eac67723a"},"cell_type":"markdown","source":"Ok, so all we can see here is that our object stands out from the rest, but is is far away from average? Let's find out."},{"metadata":{"trusted":true,"_uuid":"89f4c1e8f06119fd37c49eddefb04ee9410daa24"},"cell_type":"code","source":"ax = analysis.light_curve_interpolate_compare(object_id=object_id, target=target)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"dc331d9c22e64b0d4234305f43a6cf12873f1e23"},"cell_type":"markdown","source":"since we have 6 passbands in our light curves, let's make the same analysis for each passband separately, so that we get a better view"},{"metadata":{"trusted":true,"_uuid":"71b184a7c87237d5d0bd6f4f4341eb9939958637"},"cell_type":"code","source":"axes = analysis.light_curve_split(object_id=object_id)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"b929b784f73bdf7a34d4106fdfd86971484fa671"},"cell_type":"code","source":"axes = analysis.light_curve_split_compared(object_id=object_id, target=target)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"e863037a39f0b28d3400e4dc04c0a3976d4bc240"},"cell_type":"code","source":"axes = analysis.light_curve_interpolate_split_compare(object_id=object_id, target=target)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"27511cd6ce3ab2e51f6c1e7eb8ba0ccc2196b189"},"cell_type":"markdown","source":"So, clearly we can see that in each passband our object stands out from the average. \n\nLet's check the sky position and galactic position of this object (ra and decl)"},{"metadata":{"trusted":true,"_uuid":"bdc73b45da45e4a4f865013ffa66230b67e7155f"},"cell_type":"code","source":"ax = analysis.sky_pos(object_id=object_id)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"f543790d904e58694e5665d641ac9a49a0bddbbe"},"cell_type":"markdown","source":"Notice the white 'X' in the bottom right part of the chart, that's where our object lies on the sky map. Let's see where on the sky map are objects from our target"},{"metadata":{"trusted":true,"_uuid":"2325f6c37964cdeb3a1c80e9524b871050c4e2af"},"cell_type":"code","source":"ax = analysis.sky_pos_target(object_id=object_id, target=target)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"c1e20f9b735acfee39449a7d07c86a4d876b9606"},"cell_type":"markdown","source":"They seem to be distributed equally. Let's do the same for galactic coordinates (gal_l and gal_b)"},{"metadata":{"trusted":true,"_uuid":"c561195067949a14cfdec3697b674e64856d7944"},"cell_type":"code","source":"ax = analysis.galactic_pos(object_id=object_id)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"5867055dd78ac4c1663df1f952de9ebf65a212ab"},"cell_type":"code","source":"ax = analysis.galactic_pos_target(object_id=object_id, target=target)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"af6c475c1f7549a295215372eebfd01f49f3a5c5"},"cell_type":"markdown","source":"As an extra, there are some other metrics, not related to specific objects\n\nFirst, plot each target's each passband average light curve"},{"metadata":{"scrolled":false,"trusted":true,"_uuid":"d0550f449524de10fa8a869421dca89807ef8975"},"cell_type":"code","source":"ax = analysis.passband_target_interpolate_matrix()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"5a89bf64d4d704c263d0b577babe826ad2b752cf"},"cell_type":"markdown","source":"Plot for each target from outside of our galaxy combined light curves"},{"metadata":{"trusted":true,"_uuid":"b385352e68cf3063c2976640cb99cb790da01c2c"},"cell_type":"code","source":"ax = analysis.extragalactic_light_curves_split()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"14629fb589a89fec17ce104f752ef9900508849e"},"cell_type":"markdown","source":"And do the same for targets from our galaxy"},{"metadata":{"trusted":true,"_uuid":"004bbf8e349cdca64dfe6d5fb3aa91dadbdca185"},"cell_type":"code","source":"ax = analysis.galactic_light_curves_split()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"6719de7de04b04d3b4296459db3180b612e5b324"},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}