diff --git a/cf_notebooks/1.PartI-DifferentiableForwardModel.ipynb b/cf_notebooks/1.PartI-DifferentiableForwardModel.ipynb index c30ad28..5f1b2c0 100644 --- a/cf_notebooks/1.PartI-DifferentiableForwardModel.ipynb +++ b/cf_notebooks/1.PartI-DifferentiableForwardModel.ipynb @@ -1,15 +1,5 @@ { "cells": [ - { - "cell_type": "markdown", - "metadata": { - "id": "view-in-github", - "colab_type": "text" - }, - "source": [ - "\"Open" - ] - }, { "cell_type": "markdown", "metadata": { @@ -1084,6 +1074,31 @@ " return {'image':x[...,0], 'psf':example['psf']}" ] }, + { + "cell_type": "code", + "source": [ + "def clean_data(example):\n", + " im = example['image']\n", + " psf = example['psf']\n", + " return (tf.math.reduce_all(tf.math.is_finite(im)) and tf.math.reduce_all(tf.math.is_finite(psf)))" + ], + "metadata": { + "id": "fS0ysRrE3R_V" + }, + "execution_count": null, + "outputs": [] + }, + { + "cell_type": "code", + "source": [ + "clean_data(next(iter(dset_cosmos)))" + ], + "metadata": { + "id": "VhlG2wid4mI_" + }, + "execution_count": null, + "outputs": [] + }, { "cell_type": "code", "execution_count": null, @@ -1095,6 +1110,7 @@ "# Load COSMOS\n", "dset_cosmos = tfds.load(\"Cosmos/23.5\", split=tfds.Split.TRAIN,\n", " data_dir='galsim/tensorflow_datasets')\n", + "dset_cosmos = dset_cosmos.filter(clean_data)\n", "dset_cosmos = dset_cosmos.cache()\n", "dset_cosmos = dset_cosmos.repeat()\n", "dset_cosmos = dset_cosmos.map(preprocess)\n", @@ -1103,6 +1119,7 @@ "# Load HSC\n", "dset_hsc = tfds.load(\"HSC\", split=tfds.Split.TRAIN,\n", " data_dir='galsim/tensorflow_datasets')\n", + "dset_hsc = dset_hsc.filter(clean_data)\n", "dset_hsc = dset_hsc.cache()\n", "dset_hsc = dset_hsc.repeat()\n", "dset_hsc = dset_hsc.shuffle(10000)\n", @@ -1317,7 +1334,7 @@ "outputs": [], "source": [ "@jax.jit\n", - "def update(params, rng, opt_state):\n", + "def update(params, rng, opt_state, batch):\n", " loss, grads = jax.value_and_grad(loss_fn)(params, rng, batch)\n", " updates, opt_state = optimizer.update(grads, opt_state)\n", " # Apply gradient descent\n", @@ -1325,6 +1342,17 @@ " return loss, params, opt_state" ] }, + { + "cell_type": "code", + "source": [ + "import copy" + ], + "metadata": { + "id": "zi0Fqh5RwVhc" + }, + "execution_count": null, + "outputs": [] + }, { "cell_type": "code", "execution_count": null, @@ -1333,12 +1361,20 @@ }, "outputs": [], "source": [ + "last_loop = [0,0,0,0, 0]\n", "for i in range(10000):\n", " batch = next(combined_dset)\n", - " loss, params, opt_state = update(params, next(rng_seq), opt_state)\n", + " rng_seed = next(rng_seq)\n", + " loss, params, opt_state = update(params, rng_seed, opt_state, batch)\n", + " if not np.isnan(last_loop[1]):\n", + " last_last_loop = copy.deepcopy(last_loop)\n", + " last_loop = [copy.deepcopy(batch), copy.deepcopy(loss), copy.deepcopy(params), copy.deepcopy(opt_state), copy.deepcopy(rng_seed)]\n", + " else:\n", + " break\n", " losses.append(loss)\n", " if i %100 ==0:\n", - " print('step',i,loss)" + " print('step',i,loss)\n", + " # print('step',i,loss)" ] }, { @@ -1361,13 +1397,15 @@ "outputs": [], "source": [ "# Create a test dataset\n", - "dset_cosmos_test = tfds.load(\"Cosmos/23.5\",\n", - " split=tfds.Split.TEST)\n", + "dset_cosmos_test = tfds.load(\"Cosmos/23.5\", split=tfds.Split.TEST,\n", + " data_dir='galsim/tensorflow_datasets')\n", + "dset_cosmos_test = dset_cosmos.filter(clean_data)\n", "dset_cosmos_test = dset_cosmos_test.cache()\n", "dset_cosmos_test = dset_cosmos_test.repeat()\n", "\n", - "dset_hsc = tfds.load(\"HSC\",\n", - " split=tfds.Split.TRAIN)\n", + "dset_hsc = tfds.load(\"HSC\", split=tfds.Split.TRAIN,\n", + " data_dir='galsim/tensorflow_datasets')\n", + "dset_hsc = dset_hsc.filter(clean_data)\n", "dset_hsc = dset_hsc.cache()\n", "dset_hsc = dset_hsc.repeat()\n", "dset_hsc = dset_hsc.shuffle(10000)\n", @@ -1438,7 +1476,6 @@ "metadata": { "kernelspec": { "display_name": "Python 3", - "language": "python", "name": "python3" }, "language_info": { @@ -1455,9 +1492,10 @@ }, "colab": { "provenance": [], - "include_colab_link": true - } + "gpuType": "V100" + }, + "accelerator": "GPU" }, "nbformat": 4, "nbformat_minor": 0 -} \ No newline at end of file +}