Skip to content
Open
Show file tree
Hide file tree
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
40 changes: 35 additions & 5 deletions gerbilizer/assess.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from gerbilizer.calibration import CalibrationAccumulator
from gerbilizer.outputs.base import ModelOutput, ProbabilisticOutput, Unit
from gerbilizer.training.configs import build_config
from gerbilizer.training.dataloaders import GerbilVocalizationDataset
from gerbilizer.training.dataloaders import GerbilConcatDataset
from gerbilizer.training.models import build_model
from gerbilizer.util import make_xy_grids, subplots

Expand Down Expand Up @@ -249,8 +249,23 @@ def assess_model(
parser.add_argument(
"--data",
type=str,
required=True,
help="Path to an h5 file on which to assess the model",
nargs='+',
help="Path(s) to h5 file(s) on which to evaluate the model.",
)

parser.add_argument(
"--proportions",
type=float,
required=False,
nargs='+',
help="Optional, proportions of each dataset given to evaluate on. If not provided, will use all of each provided dataset.",
)

parser.add_argument(
"--data_random_seed",
type=int,
default=2023,
help="Optional, random seed used to select data subsets, if arg `proportions` is provided.",
)

parser.add_argument(
Expand Down Expand Up @@ -288,6 +303,16 @@ def assess_model(
f"Requested config JSON file could not be found: {args.config}"
)

# if proportions isn't provided, by default use all of each provided dataset
if args.proportions is None:
args.proportions = [1. for _ in range(len(args.data))]
elif len(args.proportions) != len(args.data):
raise ValueError(
'Must provide one proportion value per dataset in kwarg `data`! '
f'Instead encountered {len(args.proportions)} proportions and '
f'{len(args.data)} datasets. '
)

config_data = build_config(args.config)

model, _ = build_model(config_data)
Expand All @@ -309,8 +334,13 @@ def assess_model(
if arena_dims_units == "CM":
arena_dims = np.array(arena_dims) * 10

dataset = GerbilVocalizationDataset(
str(args.data),
make_xcorrs = config_data["DATA"]["COMPUTE_XCORRS"]
crop_length = config_data["DATA"]["CROP_LENGTH"]

dataset = GerbilConcatDataset(
datapaths=args.data,
proportions=args.proportions,
selection_random_seed=args.data_random_seed,
arena_dims=arena_dims,
make_xcorrs=config_data["DATA"]["COMPUTE_XCORRS"],
crop_length=config_data["DATA"]["CROP_LENGTH"],
Expand Down
45 changes: 41 additions & 4 deletions gerbilizer/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import os
import pathlib
import time

from os import path

import h5py
Expand All @@ -10,6 +11,7 @@
from gerbilizer.outputs.base import Unit
from gerbilizer.training.configs import build_config
from gerbilizer.training.trainer import Trainer
from gerbilizer.training.dataloaders import build_multi_source_datasets


def get_args():
Expand All @@ -28,7 +30,23 @@ def get_args():
parser.add_argument(
"--data",
type=str,
help="Path to directory containing train, test, and validation datasets or single h5 file for inference",
nargs='+',
help="Path to directory (or multiple directories) containing train, test, and validation datasets or one single h5 file for inference",
)

parser.add_argument(
"--proportions",
type=float,
required=False,
nargs='+',
help="Optional, proportions of each dataset given to train on. If not provided, will use all of each provided dataset.",
)

parser.add_argument(
"--data_random_seed",
type=int,
default=2023,
help="Optional, random seed used to select data subsets, if arg `proportions` is provided.",
)

parser.add_argument(
Expand Down Expand Up @@ -100,6 +118,16 @@ def validate_args(args):
else:
args.data = args.config_data["DATA"]["DATAFILE_PATH"]

# if proportions isn't provided, by default use all of each provided dataset
if args.proportions is None:
args.proportions = [1. for _ in range(len(args.data))]
elif len(args.proportions) != len(args.data):
raise ValueError(
'Must provide one proportion value per dataset in kwarg `data`! '
f'Instead encountered {len(args.proportions)} proportions and '
f'{len(args.data)} datasets. '
)

args.job_id = next_available_job_id(
args.config_data["GENERAL"]["CONFIG_NAME"], args.save_path
)
Expand All @@ -123,11 +151,11 @@ def validate_args(args):
def run_eval(args: argparse.Namespace, trainer: Trainer):
# expects args.data to point toward a file rather than a directory
# In this case, all three h5py.File objects held by the Trainer are None
data_path = args.data
data_path = args.data[0] # args.data is always a list bc of the `nargs` flag
arena_dims: tuple[float, float] = args.config_data["DATA"]["ARENA_DIMS"]
if not (data_path.endswith(".h5") or data_path.endswith(".hdf5")):
raise ValueError(
"--data argument should point to an HDF5 file with .h5 or .hdf5 file extension"
"In eval mode, --data argument should point to an HDF5 file with .h5 or .hdf5 file extension"
)
if args.output_path is not None:
dest_path = args.output_path
Expand Down Expand Up @@ -173,11 +201,20 @@ def run_eval(args: argparse.Namespace, trainer: Trainer):
def run(args):
weights = args.config_data.get("WEIGHTS_PATH", None)

train_set, val_set, test_set = build_multi_source_datasets(
config=args.config_data,
data_dirs=args.data,
proportions=args.proportions,
selection_random_seed=args.data_random_seed
)

# This modifies args.config_data['WEIGHTS_PATH']
trainer = Trainer(
data_dir=args.data,
model_dir=args.model_dir,
config_data=args.config_data,
train_set=train_set,
val_set=val_set,
test_set=test_set,
eval_mode=args.eval,
)

Expand Down
164 changes: 161 additions & 3 deletions gerbilizer/training/dataloaders.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,9 @@
"""

import os
import pathlib
import typing

from itertools import combinations
from math import comb
from typing import Optional, Tuple, Union
Expand All @@ -13,7 +16,7 @@
from scipy.signal import correlate
from torch.nn import functional as F
from torch.nn.utils.rnn import pad_sequence
from torch.utils.data import DataLoader, Dataset
from torch.utils.data import DataLoader, Dataset, ConcatDataset, Subset


class GerbilVocalizationDataset(Dataset):
Expand All @@ -34,10 +37,11 @@ def __init__(
inference (bool, optional): When true, labels will be returned in addition to data. Defaults to False.
crop_length (int): When provided, will serve random crops of fixed length instead of full vocalizations
"""
if isinstance(datapath, str):
if isinstance(datapath, str) or isinstance(datapath, pathlib.Path):
self.datapath = datapath
self.dataset = h5py.File(datapath, "r")
else:
self.dataset = datapath
raise ValueError('Expected arg `datapath` to be a str or Path pointing to an HDF5 dataset!')

if "len_idx" not in self.dataset:
raise ValueError("Improperly formatted dataset")
Expand All @@ -47,6 +51,9 @@ def __init__(
self.arena_dims = arena_dims
self.crop_length = crop_length
self.n_channels = None

def __str__(self):
return f'<GerbilVocalizationDataset object from path {self.datapath}>'

def __len__(self):
return len(self.dataset['len_idx']) - 1
Expand Down Expand Up @@ -183,6 +190,157 @@ def __processed_data_for_index__(self, idx: int):
return sound
return sound, location

class GerbilConcatDataset(ConcatDataset):
def __init__(
self,
datapaths: list[str],
proportions: list[float],
selection_random_seed: int = 2023,
*,
make_xcorrs: bool = False,
inference: bool = False,
crop_length: int = 8192,
arena_dims: Optional[Union[np.ndarray, Tuple[float, float]]] = None,
):
"""Utility class to help concatenate subsets of GerbilVocalizationDataset objects.

Args:
datapaths: Paths to HDF5 representations of GerbilVocalizationDataset objects
proportions: Indicates what size subset to select from each constituent dataset.
selection_random_seed: Seed used to randomly select subsets of each constituent dataset.
make_xcorrs (bool, optional): Triggers computation of pairwise correlations between input channels. Defaults to False.
inference (bool, optional): When true, labels will be returned in addition to data. Defaults to False.
crop_length (int): When provided, will serve random crops of fixed length instead of full vocalizations
"""
self.datapaths = datapaths
self.proportions = proportions
self.selection_random_seed = selection_random_seed

full_datasets = [
GerbilVocalizationDataset(
path,
arena_dims=arena_dims,
make_xcorrs=make_xcorrs,
crop_length=crop_length,
inference=inference
) for path in datapaths
]
# sample the subsets, storing indices to test reproducibility
rng = np.random.default_rng(seed=selection_random_seed)
self.subset_indices = []
subsets = []
for dataset, proportion in zip(full_datasets, proportions):
n_to_choose = int(proportion * len(dataset))
indices = rng.choice(len(dataset), size=n_to_choose, replace=False).tolist()
self.subset_indices.append(indices)
# and create the Subset
subsets.append(Subset(dataset, indices))
super().__init__(subsets)

@property
def n_vocalizations(self):
"""
The number of vocalizations contained in this Dataset object.
"""
return len(self)

def __str__(self):
display_strs = [
f"{prop * 100}% of data from path {datapath}"
for prop, datapath in zip(self.proportions, self.datapaths)
]
# human readability!
if len(display_strs) > 1:
display_strs[-1] = 'and ' + display_strs[-1]
display = ", ".join(display_strs)
return (f"<GerbilConcatDataset object, containing {display}, "
f"with random seed {self.selection_random_seed}>")


def build_single_source_datasets(
config: dict,
data_dir: str
) -> tuple[GerbilVocalizationDataset, GerbilVocalizationDataset, GerbilVocalizationDataset]:
"""
Construct three GerbilVocalizationDataset objects (train, val, test) from
files in a single source directory `data_dir`, which should point to directory
containing files `train_set.h5`, `val_set.h5` and `test_set.h5`.

Example usage:
```
>>> from gerbilizer.training.configs import build_config
>>> config = build_config('/path/to/model_config.json')
>>> data_dir = '/path/to/data/set/directory'
>>> train, val, test = build_single_source_datasets(config, data_dir)
```
"""
arena_dims = config["DATA"]["ARENA_DIMS"]
make_xcorrs = config["DATA"]["COMPUTE_XCORRS"]
crop_length = config["DATA"]["CROP_LENGTH"]

data_filenames = ("train_set.h5", "val_set.h5", "test_set.h5")
use_inference_flags = (False, False, True)

datasets = []
for filename, use_inference in zip(data_filenames, use_inference_flags):
datasets.append(GerbilVocalizationDataset(
os.path.join(data_dir, filename),
arena_dims=arena_dims,
make_xcorrs=make_xcorrs,
crop_length=crop_length,
inference=use_inference
))

return tuple(datasets)

def build_multi_source_datasets(
config: dict,
data_dirs: list[str],
proportions: list[float],
selection_random_seed: int = 2023
) -> tuple[GerbilConcatDataset, GerbilConcatDataset, GerbilConcatDataset]:
"""
Construct three `GerbilConcatDataset` objects (train, val, test) from
files in a list of source directories `data_dirs`.

Args:
config: Loaded model configuration dictionary.
data_dirs: A list of directories, with each directory containing `train_set.h5`, `val_set.h5`, and `test_set.h5`.
proportions: Floats indicating the proportion of each source to include.
selection_random_seed: Random seed used to select subsets of each dataset.

Example usage:
```
>>> from gerbilizer.training.configs import build_config
>>> config = build_config('/path/to/model_config.json')
>>> data_dirs = ['/data/set/number/one', '/data/set/number/two']
>>> # include 90% of dataset 1, 30% of dataset 2
>>> proportions = [0.9, 0.3]
>>> train, val, test = build_multi_source_datasets(config, data_dirs, proportions)
```
"""
arena_dims = config["DATA"]["ARENA_DIMS"]
make_xcorrs = config["DATA"]["COMPUTE_XCORRS"]
crop_length = config["DATA"]["CROP_LENGTH"]

data_filenames = ("train_set.h5", "val_set.h5", "test_set.h5")
use_inference_flags = (False, False, True)

datasets = []
for filename, use_inference in zip(data_filenames, use_inference_flags):
datapaths = [os.path.join(dirname, filename) for dirname in data_dirs]
datasets.append(GerbilConcatDataset(
datapaths=datapaths,
proportions=proportions,
selection_random_seed=selection_random_seed,
arena_dims=arena_dims,
make_xcorrs=make_xcorrs,
crop_length=crop_length,
inference=use_inference
))

return tuple(datasets)


def build_dataloaders(path_to_data: str, config: dict):
# Construct Dataset objects.
Expand Down
Loading