mirror of
https://github.com/google-deepmind/deepmind-research.git
synced 2026-10-06 06:59:41 +08:00
Initial release of A Deep Learning Approach for Characterizing Major Galaxy Mergers
PiperOrigin-RevId: 369646863
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
```
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user