Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 59 additions & 21 deletions cf_notebooks/1.PartI-DifferentiableForwardModel.ipynb
Original file line number Diff line number Diff line change
@@ -1,15 +1,5 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {
"id": "view-in-github",
"colab_type": "text"
},
"source": [
"<a href=\"https://colab.research.google.com/github/EiffL/Quarks2CosmosDataChallenge/blob/colab/notebooks/PartI-DifferentiableForwardModel.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>"
]
},
{
"cell_type": "markdown",
"metadata": {
Expand Down Expand Up @@ -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,
Expand All @@ -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",
Expand All @@ -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",
Expand Down Expand Up @@ -1317,14 +1334,25 @@
"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",
" params = optax.apply_updates(params, updates)\n",
" return loss, params, opt_state"
]
},
{
"cell_type": "code",
"source": [
"import copy"
],
"metadata": {
"id": "zi0Fqh5RwVhc"
},
"execution_count": null,
"outputs": []
},
{
"cell_type": "code",
"execution_count": null,
Expand All @@ -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)"
]
},
{
Expand All @@ -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",
Expand Down Expand Up @@ -1438,7 +1476,6 @@
"metadata": {
"kernelspec": {
"display_name": "Python 3",
"language": "python",
"name": "python3"
},
"language_info": {
Expand All @@ -1455,9 +1492,10 @@
},
"colab": {
"provenance": [],
"include_colab_link": true
}
"gpuType": "V100"
},
"accelerator": "GPU"
},
"nbformat": 4,
"nbformat_minor": 0
}
}