import collections import csv import json import logging import os import pathlib import random from collections import defaultdict from pathlib import Path from pprint import pformat import numpy as np import torch import torchvision.datasets as datasets from .common import ImageFolderWithPaths, SubsetSampler from .imagenet import ImageNet def load_labels(input_paths): if isinstance(input_paths, (Path, str)): input_paths = [input_paths] # Map data to list of (input_path, label) tuples. labels = collections.defaultdict(list) labels_list = None for path in input_paths: with open(path, 'r') as f: annotations = json.load(f) for i, row in enumerate(annotations['annotations']): labels[row['key']].append((path, row)) if labels_list is None: labels_list = annotations['labels'] else: assert labels_list == annotations['labels'] for key, key_labels in labels.items(): if len(key_labels) > 1: paths = [x[0] for x in key_labels] logging.debug( f'{key} labeled multiple times in {paths}; using latest ' f'label from {paths[-1]}.') labels[key] = key_labels[-1][1] return labels, labels_list def filter_labels(labels, labels_list, file_logger=None, must_have=[], must_not_have=[], can_have=[], must_have_one_of=False, unspecified_labels_policy='error', return_nonmatching=False): if file_logger is None: file_logger = logging.getLogger() label_map = {} label_names = {} for i, label in enumerate(labels_list): label_map[label] = i label_names[i] = label def validate_label(label): if label not in label_map: raise ValueError('Unknown label %s, valid labels: %s' % (label, label_map.keys())) return True must_have_labels = set( [label_map[x] for x in must_have if validate_label(x)]) must_not_have_labels = set( [label_map[x] for x in must_not_have if validate_label(x)]) can_have_labels = set( [label_map[x] for x in can_have if validate_label(x)]) unspecified_labels = ( set(label_map.values()) - (must_have_labels | must_not_have_labels | can_have_labels)) if unspecified_labels: if unspecified_labels_policy == 'error': raise ValueError('Label(s): %s were not specified in any of ' '--{must,must-not,can}-have.' % [label_names[x] for x in unspecified_labels]) elif unspecified_labels_policy == 'can-have': can_have_labels |= unspecified_labels elif unspecified_labels_policy == 'must-have': must_have_labels |= unspecified_labels elif unspecified_labels_policy == 'must-not-have': must_not_have_labels |= unspecified_labels else: raise ValueError('Unknown unspecified_labels_policy %s' % unspecified_labels_policy) logging.info('Looking for rows that') if must_have_one_of: logging.info('MUST HAVE (one of): %s', [label_names[x] for x in must_have_labels]) else: logging.info('MUST HAVE: %s', [label_names[x] for x in must_have_labels]) logging.info('MUST NOT HAVE: %s', [label_names[x] for x in must_not_have_labels]) logging.info('CAN HAVE: %s', [label_names[x] for x in can_have_labels]) valid_rows = [] invalid_rows = [] for key, row in labels.items(): row_labels = set(row['labels']) missing_labels = must_have_labels - row_labels unwanted_labels = row_labels & must_not_have_labels if must_have_one_of: if missing_labels == must_have_labels: file_logger.info('Label %s missing labels %s' % (pformat( dict(row)), [label_names[x] for x in missing_labels])) invalid_rows.append(row) continue elif missing_labels: file_logger.info( 'Label %s missing labels %s' % (pformat(dict(row)), [label_names[x] for x in missing_labels])) invalid_rows.append(row) continue if unwanted_labels: file_logger.info( 'Label %s has unwanted labels %s' % (pformat(dict(row)), [label_names[x] for x in unwanted_labels])) invalid_rows.append(row) continue valid_rows.append(row) if return_nonmatching: return valid_rows, invalid_rows else: return valid_rows def evaluate_pmk(predictions, labels, valid_pmk): """ Args: predictions (Dict[str, np.array]) labels (Dict[str, List[int]]): Labels for anchor frames. valid_pmk (Dict[str, Dict[int, str]]) """ anchor_is_correct = {} pmk_is_correct = {} for anchor, pmk_dict in valid_pmk.items(): anchor_labels = labels[anchor] anchor_prediction = predictions[anchor].argmax() anchor_is_correct[anchor] = anchor_prediction in anchor_labels pmk_is_correct[anchor] = {} for offset, pmk_key in pmk_dict.items(): pmk_prediction = predictions[pmk_key].argmax() pmk_is_correct[anchor][pmk_key] = pmk_prediction in anchor_labels return anchor_is_correct, pmk_is_correct def create_pmk_score(predictions_by_key, anchor_labels, pmk_frames): """ Args: predictions_by_key (Dict[str, np.array]) anchor_labels (Dict[str, List[int]]): Labels for anchor frames. pmk_frames (Dict[str, Dict[int, str]]): Map anchor frame to dict mapping valid pmk offset to pmk frame key. """ pmk_frames = pmk_frames.copy() for anchor in anchor_labels: if anchor not in pmk_frames: pmk_frames[anchor] = {} anchor_is_correct, pmk_is_correct = evaluate_pmk(predictions_by_key, anchor_labels, pmk_frames) correct_anchors = { k for k, correct in anchor_is_correct.items() if correct } all_anchors = [k for k, correct in anchor_is_correct.items()] num_anchor_correct = len(correct_anchors) anchor_accuracy = num_anchor_correct / max(len(anchor_is_correct), 1e-9) pmk_correct = [ anchor for anchor in correct_anchors if all(pmk_is_correct[anchor].values()) ] rand_correct = [ anchor for anchor in all_anchors if (len(pmk_is_correct[anchor].values()) == 0) or random.choice(list(pmk_is_correct[anchor].values())) ] pmk_accuracy = len(pmk_correct) / max(len(anchor_is_correct), 1e-9) rand_accuracy = len(rand_correct) / max(len(anchor_is_correct), 1e-9) # Collect auxiliary data. benign_frames = sorted(anchor_labels.keys()) adversarial_pmk = {} # Map anchor to list of adversarial pmk offsets nonadversarial_pmk = {} # Map anchor to list of non-adv pmk offsets for anchor in benign_frames: if not anchor_is_correct[anchor]: adversarial_pmk[anchor] = None nonadversarial_pmk[anchor] = None else: incorrect_frames = [] correct_frames = [] for i, (offset, pmk_key) in enumerate(pmk_frames[anchor].items()): if pmk_is_correct[anchor][pmk_key]: correct_frames.append(offset) else: incorrect_frames.append(offset) adversarial_pmk[anchor] = incorrect_frames nonadversarial_pmk[anchor] = correct_frames score_info = {} score_info["benign_accuracy"] = anchor_accuracy score_info["benign_frames"] = benign_frames score_info["adversarial_pmk"] = adversarial_pmk score_info["nonadversarial_pmk"] = nonadversarial_pmk score_info["pmk_keys"] = pmk_frames score_info["correct_anchors"] = sorted(correct_anchors) score_info["incorrect_anchors"] = sorted( set(anchor_is_correct.keys()) - correct_anchors) score_info["l_infs"] = [] # TODO return pmk_accuracy, score_info def ms_to_frame_15fps(ms): return round(ms / 1000 * 15) def path_to_key(path): path = pathlib.Path(path) return f"{path.parent.name}/{path.name}" def get_pmk_key(anchor_key, pmk_index): """Returns pmk portion of pmk key. The full pm-k key, as used in annotations, is '{anchor_key},{pmk_key}'.""" video, anchor_index, anchor_ms = parse_frame_key(anchor_key) prefix = f'{video}_{anchor_ms}' return f'{prefix}/frame-{pmk_index}.jpg' def split_pmk_key(key): anchor_path, pmk_path = key.split(",") return path_to_key(anchor_path), path_to_key(pmk_path) def parse_frame_key(key, return_ms=True): """Parse key into video, frame index, and anchor ms.""" key = path_to_key(key) # Key format: