| import json |
| import os |
| import pathlib |
|
|
| import numpy as np |
| import pandas as pd |
| import wilds |
| from wilds.common.data_loaders import get_eval_loader, get_train_loader |
| from wilds.datasets.wilds_dataset import WILDSSubset |
|
|
|
|
| def get_mask_non_empty(dataset): |
| metadf = pd.read_csv(dataset._data_dir / 'metadata.csv') |
| filename = os.path.expanduser(dataset._data_dir / 'iwildcam2020_megadetector_results.json') |
| with open(filename, 'r') as f: |
| md_data = json.load(f) |
| id_to_maxdet = {x['id']: x['max_detection_conf'] for x in md_data['images']} |
| threshold = 0.95 |
| mask_non_empty = [id_to_maxdet[x] >= threshold for x in metadf['image_id']] |
| return mask_non_empty |
|
|
|
|
| def get_nonempty_subset(dataset, split, frac=1.0, transform=None): |
| if split not in dataset.split_dict: |
| raise ValueError(f"Split {split} not found in dataset's split_dict.") |
| split_mask = dataset.split_array == dataset.split_dict[split] |
|
|
| |
| |
| mask_non_empty = get_mask_non_empty(dataset) |
| split_mask = split_mask & mask_non_empty |
|
|
| split_idx = np.where(split_mask)[0] |
| if frac < 1.0: |
| num_to_retain = int(np.round(float(len(split_idx)) * frac)) |
| split_idx = np.sort(np.random.permutation(split_idx)[:num_to_retain]) |
| subset = WILDSSubset(dataset, split_idx, transform) |
| return subset |
|
|
|
|
| class IWildCam: |
| def __init__(self, |
| preprocess, |
| location=os.path.expanduser('~/data'), |
| remove_non_empty=False, |
| batch_size=128, |
| num_workers=16, |
| classnames=None, |
| subset='train'): |
| self.dataset = wilds.get_dataset(dataset='iwildcam', root_dir=location) |
| self.train_dataset = self.dataset.get_subset('train', transform=preprocess) |
| self.train_loader = get_train_loader("standard", self.train_dataset, num_workers=num_workers, batch_size=batch_size) |
|
|
| if remove_non_empty: |
| self.train_dataset = get_nonempty_subset(self.dataset, 'train', transform=preprocess) |
| else: |
| self.train_dataset = self.dataset.get_subset('train', transform=preprocess) |
|
|
| if remove_non_empty: |
| self.test_dataset = get_nonempty_subset(self.dataset, subset, transform=preprocess) |
| else: |
| self.test_dataset = self.dataset.get_subset(subset, transform=preprocess) |
|
|
| self.test_loader = get_eval_loader( |
| "standard", self.test_dataset, |
| num_workers=num_workers, |
| batch_size=batch_size) |
|
|
| labels_csv = pathlib.Path(__file__).parent / 'iwildcam_metadata' / 'labels.csv' |
| df = pd.read_csv(labels_csv) |
| df = df[df['y'] < 99999] |
| |
| self.classnames = [s.lower() for s in list(df['english'])] |
|
|
| def post_loop_metrics(self, labels, preds, metadata, args): |
| preds = preds.argmax(dim=1, keepdim=True).view_as(labels) |
| results = self.dataset.eval(preds, labels, metadata) |
| return results[0] |
|
|
|
|
| class IWildCamID(IWildCam): |
| def __init__(self, *args, **kwargs): |
| kwargs['subset'] = 'id_test' |
| super().__init__(*args, **kwargs) |
|
|
|
|
| class IWildCamOOD(IWildCam): |
| def __init__(self, *args, **kwargs): |
| kwargs['subset'] = 'test' |
| super().__init__(*args, **kwargs) |
|
|
|
|
| class IWildCamNonEmpty(IWildCam): |
| def __init__(self, *args, **kwargs): |
| kwargs['subset'] = 'train' |
| super().__init__(*args, **kwargs) |
|
|
|
|
| class IWildCamIDNonEmpty(IWildCam): |
| def __init__(self, *args, **kwargs): |
| kwargs['subset'] = 'id_test' |
| super().__init__(*args, **kwargs) |
|
|
|
|
| class IWildCamOODNonEmpty(IWildCam): |
| def __init__(self, *args, **kwargs): |
| kwargs['subset'] = 'test' |
| super().__init__(*args, **kwargs) |
|
|