Initial release of A Deep Learning Approach for Characterizing Major Galaxy Mergers

PiperOrigin-RevId: 369646863
This commit is contained in:
Louise Deason
2021-04-28 14:53:08 +00:00
parent fe4a129143
commit 3dc0baece1
12 changed files with 1538 additions and 0 deletions
+1
View File
@@ -24,6 +24,7 @@ https://deepmind.com/research/publications/
## Projects
* [A Deep Learning Approach for Characterizing Major Galaxy Mergers](galaxy_mergers)
* [Better, Faster Fermionic Neural Networks](kfac_ferminet_alpha) (KFAC implementation)
* [Object-based attention for spatio-temporal reasoning](object_attention_for_reasoning)
* [Effective gene expression prediction from sequence by integrating long-range interactions](enformer)
+41
View File
@@ -0,0 +1,41 @@
# A Deep Learning Approach for Characterizing Major Galaxy Mergers
This repository contains evaluation code and checkpoints to reproduce
figures in https://arxiv.org/abs/2102.05182.
The main evaluation module is `main.py`. It uses the provided checkpoint path
and dataset path to run evaluation.
## Setup
To set up a Python virtual environment with the required dependencies, run:
```shell
python3 -m venv galaxy_mergers_env
source galaxy_mergers_env/bin/activate
pip install --upgrade pip setuptools wheel
pip install -r requirements.txt
```
### License
While the code is licensed under the Apache 2.0 License, the checkpoints weights
are made available for non-commercial use only under the terms of the
Creative Commons Attribution-NonCommercial 4.0 International (CC BY-NC 4.0)
license. You can find details at:
https://creativecommons.org/licenses/by-nc/4.0/legalcode.
### Citing our work
If you use this work, consider citing our paper:
```bibtex
@article{koppula2021deep,
title={A Deep Learning Approach for Characterizing Major Galaxy Mergers},
author={Koppula, Skanda and Bapst, Victor and Huertas-Company, Marc and Blackwell, Sam and Grabska-Barwinska, Agnieszka and Dieleman, Sander and Huber, Andrea and Antropova, Natasha and Binkowski, Mikolaj and Openshaw, Hannah and others},
journal={Workshop for Machine Learning and the Physical Sciences @ NeurIPS 2020},
year={2021}
}
```
+87
View File
@@ -0,0 +1,87 @@
# Copyright 2021 DeepMind Technologies Limited.
#
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Helpers to pre-process Antennae galaxy images."""
import collections
import os
from astropy.io import fits
import numpy as np
from scipy import ndimage
import tensorflow.compat.v2 as tf
def norm_antennae_images(images, scale=1000):
return tf.math.asinh(images/scale)
def renorm_antennae(images):
median = np.percentile(images.numpy().flatten(), 50)
img_range = np.ptp(images.numpy().flatten())
return (images - median) / (img_range / 2)
def get_antennae_images(antennae_fits_dir):
"""Load the raw Antennae galaxy images."""
all_fits_files = [
os.path.join(antennae_fits_dir, f)
for f in os.listdir(antennae_fits_dir)
]
freq_mapping = {'red': 160, 'blue': 850}
paired_fits_files = collections.defaultdict(list)
for f in all_fits_files:
redshift = float(f[-8:-5])
paired_fits_files[redshift].append(f)
for redshift, files in paired_fits_files.items():
paired_fits_files[redshift] = sorted(
files, key=lambda f: freq_mapping[f.split('/')[-1].split('_')[0]])
print('Reading files:', paired_fits_files)
print('Redshifts:', sorted(paired_fits_files.keys()))
galaxy_views = collections.defaultdict(list)
for redshift in paired_fits_files:
for view_path in paired_fits_files[redshift]:
with open(view_path, 'rb') as f:
fits_data = fits.open(f)
galaxy_views[redshift].append(np.array(fits_data[0].data))
batched_images = []
for redshift in paired_fits_files:
img = tf.constant(np.array(galaxy_views[redshift]))
img = tf.transpose(img, (1, 2, 0))
img = tf.image.resize(img, size=(60, 60))
batched_images.append(img)
return tf.stack(batched_images)
def preprocess_antennae_images(antennae_images):
"""Pre-process the Antennae galaxy images into a reasonable range."""
rotated_antennae_images = [
ndimage.rotate(img, 10, reshape=True, cval=-1)[10:-10, 10:-10]
for img in antennae_images
]
rotated_antennae_images = [
np.clip(img, 0, 1e9) for img in rotated_antennae_images
]
rotated_antennae_images = tf.stack(rotated_antennae_images)
normed_antennae_images = norm_antennae_images(rotated_antennae_images)
normed_antennae_images = tf.clip_by_value(normed_antennae_images, 1, 4.5)
renormed_antennae_images = renorm_antennae(normed_antennae_images)
return renormed_antennae_images
+98
View File
@@ -0,0 +1,98 @@
# Copyright 2021 DeepMind Technologies Limited.
#
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Default config, focused on model evaluation."""
from ml_collections import config_dict
def get_config(filter_time_intervals=None):
"""Return config object for training."""
config = config_dict.ConfigDict()
config.eval_strategy = config_dict.ConfigDict()
config.eval_strategy.class_name = 'OneDeviceConfig'
config.eval_strategy.kwargs = config_dict.ConfigDict(
dict(device_type='v100'))
## Experiment config.
config.experiment_kwargs = config_dict.ConfigDict(dict(
resnet_kwargs=dict(
blocks_per_group_list=[3, 4, 6, 3], # This choice is ResNet50.
bn_config=dict(
decay_rate=0.9,
eps=1e-5),
resnet_v2=False,
additional_features_mode='mlp',
),
optimizer_config=dict(
class_name='Momentum',
kwargs={'momentum': 0.9},
# Set up the learning rate schedule.
lr_init=0.025,
lr_factor=0.1,
lr_schedule=(50e3, 100e3, 150e3),
gradient_clip=5.,
),
l2_regularization=1e-4,
total_train_batch_size=128,
train_net_args={'is_training': True},
eval_batch_size=128,
eval_net_args={'is_training': True},
data_config=dict(
# dataset loading
dataset_path=None,
num_val_splits=10,
val_split=0,
# image cropping
image_size=(80, 80, 7),
train_crop_type='crop_fixed',
test_crop_type='crop_fixed',
n_crop_repeat=1,
train_augmentations=dict(
rotation_and_flip=True,
rescaling=True,
translation=True,
),
test_augmentations=dict(
rotation_and_flip=False,
rescaling=False,
translation=False,
),
test_time_ensembling='sum',
num_eval_buckets=5,
eval_confidence_interval=95,
task='grounded_unnormalized_regression',
loss_config=dict(
loss='mse',
mse_normalize=False,
),
model_uncertainty=True,
additional_features='',
time_filter_intervals=filter_time_intervals,
class_boundaries={
'0': [[-1., 0]],
'1': [[0, 1.]]
},
frequencies_to_use='all',
),
n_train_epochs=100
))
return config
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,68 @@
# Copyright 2021 DeepMind Technologies Limited.
#
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Helpers to visualize gradients and other interpretability analysis."""
import numpy as np
import tensorflow.compat.v2 as tf
def rotate_by_right_angle_multiple(image, rot=90):
"""Rotate an image by right angles."""
if rot not in [0, 90, 180, 270]:
raise ValueError(f"Cannot rotate by non-90 degree angle {rot}")
if rot in [90, -270]:
image = np.transpose(image, (1, 0, 2))
image = image[::-1]
elif rot in [180, -180]:
image = image[::-1, ::-1]
elif rot in [270, -90]:
image = np.transpose(image, (1, 0, 2))
image = image[:, ::-1]
return image
def compute_gradient(images, evaluator, is_training=False):
inputs = tf.Variable(images[None], dtype=tf.float32)
with tf.GradientTape() as tape:
tape.watch(inputs)
time_sigma = evaluator.model(inputs, None, is_training)
grad_time = tape.gradient(time_sigma[:, 0], inputs)
return grad_time, time_sigma
def compute_grads_for_rotations(images, evaluator, is_training=False):
test_gradients, test_outputs = [], []
for rotation in np.arange(0, 360, 90):
images_rot = rotate_by_right_angle_multiple(images, rotation)
grads, time_sigma = compute_gradient(images_rot, evaluator, is_training)
grads = np.squeeze(grads.numpy())
inv_grads = rotate_by_right_angle_multiple(grads, -rotation)
test_gradients.append(inv_grads)
test_outputs.append(time_sigma.numpy())
return np.squeeze(test_gradients), np.squeeze(test_outputs)
def compute_grads_for_rotations_and_flips(images, evaluator):
grads, time_sigma = compute_grads_for_rotations(images, evaluator)
grads_f, time_sigma_f = compute_grads_for_rotations(images[::-1], evaluator)
grads_f = grads_f[:, ::-1]
all_grads = np.concatenate([grads, grads_f], 0)
model_outputs = np.concatenate((time_sigma, time_sigma_f), 0)
return all_grads, model_outputs
+169
View File
@@ -0,0 +1,169 @@
# Copyright 2021 DeepMind Technologies Limited.
#
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Helpers to compute loss metrics."""
import scipy.stats
import tensorflow.compat.v2 as tf
import tensorflow_probability as tfp
TASK_CLASSIFICATION = 'classification'
TASK_NORMALIZED_REGRESSION = 'normalized_regression'
TASK_UNNORMALIZED_REGRESSION = 'unnormalized_regression'
TASK_GROUNDED_UNNORMALIZED_REGRESSION = 'grounded_unnormalized_regression'
REGRESSION_TASKS = [TASK_NORMALIZED_REGRESSION, TASK_UNNORMALIZED_REGRESSION,
TASK_GROUNDED_UNNORMALIZED_REGRESSION]
ALL_TASKS = [TASK_CLASSIFICATION] + REGRESSION_TASKS
LOSS_MSE = 'mse'
LOSS_SOFTMAX_CROSS_ENTROPY = 'softmax_cross_entropy'
ALL_LOSSES = [LOSS_SOFTMAX_CROSS_ENTROPY, LOSS_MSE]
def normalize_regression_loss(regression_loss, predictions):
# Normalize loss such that:
# 1) E_{x uniform}[loss(x, prediction)] does not depend on prediction
# 2) E_{x uniform, prediction uniform}[loss(x, prediction)] is as before.
# Divides MSE regression loss by E[(prediction-x)^2]; assumes x=[-1,1]
normalization = 2./3.
normalized_loss = regression_loss / ((1./3 + predictions**2) / normalization)
return normalized_loss
def equal32(x, y):
return tf.cast(tf.equal(x, y), tf.float32)
def mse_loss(predicted, targets):
return (predicted - targets) ** 2
def get_std_factor_from_confidence_percent(percent):
dec = percent/100.
inv_dec = 1 - dec
return scipy.stats.norm.ppf(dec+inv_dec/2)
def get_all_metric_names(task_type, model_uncertainty, loss_config, # pylint: disable=unused-argument
mode='eval', return_dict=True):
"""Get all the scalar fields produced by compute_loss_and_metrics."""
names = ['regularization_loss', 'prediction_accuracy', str(mode)+'_loss']
if task_type == TASK_CLASSIFICATION:
names += ['classification_loss']
else:
names += ['regression_loss', 'avg_mu', 'var_mu']
if model_uncertainty:
names += ['uncertainty_loss', 'scaled_regression_loss',
'uncertainty_plus_scaled_regression',
'avg_sigma', 'var_sigma',
'percent_in_conf_interval', 'error_sigma_correlation',
'avg_prob']
if return_dict:
return {name: 0. for name in names}
else:
return names
def compute_loss_and_metrics(mu, log_sigma_sq,
regression_targets, labels,
task_type, model_uncertainty, loss_config,
regularization_loss=0., confidence_interval=95,
mode='train'):
"""Computes loss statistics and other metrics."""
scalars_to_log = dict()
vectors_to_log = dict()
scalars_to_log['regularization_loss'] = regularization_loss
vectors_to_log['mu'] = mu
if task_type == TASK_CLASSIFICATION:
cross_entropy = tf.nn.sparse_softmax_cross_entropy_with_logits(
logits=mu, labels=labels, name='cross_entropy')
classification_loss = tf.reduce_mean(cross_entropy, name='class_loss')
total_loss = classification_loss
sigma = None
scalars_to_log['classification_loss'] = classification_loss
predicted_labels = tf.argmax(mu, axis=1)
correct_predictions = equal32(predicted_labels, labels)
else:
regression_loss = mse_loss(mu, regression_targets)
if 'mse_normalize' in loss_config and loss_config['mse_normalize']:
assert task_type in [TASK_GROUNDED_UNNORMALIZED_REGRESSION,
TASK_NORMALIZED_REGRESSION]
regression_loss = normalize_regression_loss(regression_loss, mu)
avg_regression_loss = tf.reduce_mean(regression_loss)
vectors_to_log['regression_loss'] = regression_loss
scalars_to_log['regression_loss'] = avg_regression_loss
scalars_to_log['avg_mu'] = tf.reduce_mean(mu)
scalars_to_log['var_mu'] = tf.reduce_mean(mse_loss(mu, tf.reduce_mean(mu)))
predicted_labels = tf.cast(mu > 0, tf.int64)
correct_predictions = equal32(predicted_labels, labels)
if model_uncertainty:
# This implements Eq. (1) in https://arxiv.org/pdf/1612.01474.pdf
inv_sigma_sq = tf.math.exp(-log_sigma_sq)
scaled_regression_loss = regression_loss * inv_sigma_sq
scaled_regression_loss = tf.reduce_mean(scaled_regression_loss)
uncertainty_loss = tf.reduce_mean(log_sigma_sq)
total_loss = uncertainty_loss + scaled_regression_loss
scalars_to_log['uncertainty_loss'] = uncertainty_loss
scalars_to_log['scaled_regression_loss'] = scaled_regression_loss
scalars_to_log['uncertainty_plus_scaled_regression'] = total_loss
sigma = tf.math.exp(log_sigma_sq / 2.)
vectors_to_log['sigma'] = sigma
scalars_to_log['avg_sigma'] = tf.reduce_mean(sigma)
var_sigma = tf.reduce_mean(mse_loss(sigma, tf.reduce_mean(sigma)))
scalars_to_log['var_sigma'] = var_sigma
# Compute # of labels that fall into the confidence interval.
std_factor = get_std_factor_from_confidence_percent(confidence_interval)
lower_bound = mu - std_factor * sigma
upper_bound = mu + std_factor * sigma
preds = tf.logical_and(tf.greater(regression_targets, lower_bound),
tf.less(regression_targets, upper_bound))
percent_in_conf_interval = tf.reduce_mean(tf.cast(preds, tf.float32))
scalars_to_log['percent_in_conf_interval'] = percent_in_conf_interval*100
error_sigma_corr = tfp.stats.correlation(x=regression_loss,
y=sigma, event_axis=None)
scalars_to_log['error_sigma_correlation'] = error_sigma_corr
dists = tfp.distributions.Normal(mu, sigma)
probs = dists.prob(regression_targets)
scalars_to_log['avg_prob'] = tf.reduce_mean(probs)
else:
total_loss = avg_regression_loss
loss_name = str(mode)+'_loss'
total_loss = tf.add(total_loss, regularization_loss, name=loss_name)
scalars_to_log[loss_name] = total_loss
vectors_to_log['correct_predictions'] = correct_predictions
scalars_to_log['prediction_accuracy'] = tf.reduce_mean(correct_predictions)
# Validate that metrics outputted are exactly what is expected
expected = get_all_metric_names(task_type, model_uncertainty,
loss_config, mode, False)
assert set(expected) == set(scalars_to_log.keys())
return scalars_to_log, vectors_to_log
+51
View File
@@ -0,0 +1,51 @@
# Copyright 2021 DeepMind Technologies Limited.
#
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Simple script to model evaluation on a checkpoint and dataset."""
import ast
from absl import app
from absl import flags
from absl import logging
from galaxy_mergers import evaluator
flags.DEFINE_string('checkpoint_path', '', 'Path to TF2 checkpoint to eval.')
flags.DEFINE_string('data_path', '', 'Path to TFRecord(s) with data.')
flags.DEFINE_string('filter_time_intervals', None,
'Merger time intervals on which to perform regression.'
'Specify None for the default time interval [-1,1], or'
' a custom list of intervals, e.g. [[-0.2,0], [0.5,1]].')
FLAGS = flags.FLAGS
def main(_) -> None:
if FLAGS.filter_time_intervals is not None:
filter_time_intervals = ast.literal_eval(FLAGS.filter_time_intervals)
else:
filter_time_intervals = None
config, ds, experiment = evaluator.get_config_dataset_evaluator(
filter_time_intervals,
FLAGS.checkpoint_path,
config_override={
'experiment_kwargs.data_config.dataset_path': FLAGS.data_path,
})
metrics, _, _ = evaluator.run_model_on_dataset(experiment, ds, config)
logging.info('Evaluation complete. Metrics: %s', metrics)
if __name__ == '__main__':
app.run(main)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+53
View File
@@ -0,0 +1,53 @@
absl-py==0.11.0
astropy==4.2
astunparse==1.6.3
cachetools==4.2.1
certifi==2020.12.5
chardet==4.0.0
cloudpickle==1.6.0
contextlib2==0.6.0.post1
cycler==0.10.0
decorator==4.4.2
dm-sonnet==2.0.0
dm-tree==0.1.5
flatbuffers==1.12
gast==0.3.3
google-auth==1.27.0
google-auth-oauthlib==0.4.2
google-pasta==0.2.0
grpcio==1.32.0
h5py==2.10.0
idna==2.10
Keras-Preprocessing==1.1.2
kiwisolver==1.3.1
Markdown==3.3.4
matplotlib==3.3.4
ml-collections==0.1.0
numpy==1.19.5
oauthlib==3.1.0
opt-einsum==3.3.0
Pillow==8.1.0
pkg-resources==0.0.0
protobuf==3.15.3
pyasn1==0.4.8
pyasn1-modules==0.2.8
pyerfa==1.7.2
pyparsing==2.4.7
python-dateutil==2.8.1
PyYAML==5.4.1
requests==2.25.1
requests-oauthlib==1.3.0
rsa==4.7.2
scipy==1.6.1
six==1.15.0
tabulate==0.8.9
tensorboard==2.4.1
tensorboard-plugin-wit==1.8.0
tensorflow==2.4.1
tensorflow-estimator==2.4.0
tensorflow-probability==0.12.1
termcolor==1.1.0
typing-extensions==3.7.4.3
urllib3==1.26.3
Werkzeug==1.0.1
wrapt==1.12.1