hard_neg: bool, add hard negative? other_neg: bool, add other negatives?(quadra loss fuction required)
(self, dict_value, hard_neg=None, other_neg=False)
| 67 | self.data_count = 0 |
| 68 | |
| 69 | def get_tuple(self, dict_value, hard_neg=None, other_neg=False): |
| 70 | ''' |
| 71 | hard_neg: bool, add hard negative? |
| 72 | other_neg: bool, add other negatives?(quadra loss fuction required) |
| 73 | ''' |
| 74 | if hard_neg is None: |
| 75 | hard_neg = [] |
| 76 | possible_negs = [] |
| 77 | |
| 78 | num_pos = self.config.TRAINING.BATCH.POSITIVES_PER_QUERY |
| 79 | num_neg = self.config.TRAINING.BATCH.NEGATIVES_PER_QUERY |
| 80 | random.shuffle(dict_value["positives"]) |
| 81 | pos_files = [] |
| 82 | for i in range(num_pos): |
| 83 | pos_files.append(self.queries[dict_value["positives"][i]]["query"]) |
| 84 | |
| 85 | neg_files = [] |
| 86 | neg_indices = [] |
| 87 | if len(hard_neg) == 0: |
| 88 | random.shuffle(dict_value["negatives"]) |
| 89 | for i in range(num_neg): |
| 90 | neg_files.append(self.queries[dict_value["negatives"][i]]["query"]) |
| 91 | neg_indices.append(dict_value["negatives"][i]) |
| 92 | else: |
| 93 | random.shuffle(dict_value["negatives"]) |
| 94 | for i in hard_neg: |
| 95 | neg_files.append(self.queries[i]["query"]) |
| 96 | neg_indices.append(i) |
| 97 | j = 0 |
| 98 | while len(neg_files) < num_neg: |
| 99 | if not dict_value["negatives"][j] in hard_neg: |
| 100 | neg_files.append(self.queries[dict_value["negatives"][j]]["query"]) |
| 101 | neg_indices.append(dict_value["negatives"][j]) |
| 102 | j += 1 |
| 103 | if other_neg: |
| 104 | # get neighbors of negatives and query |
| 105 | neighbors = [] |
| 106 | for pos in dict_value["positives"]: |
| 107 | neighbors.append(pos) |
| 108 | for neg in neg_indices: |
| 109 | for pos in self.queries[neg]["positives"]: |
| 110 | neighbors.append(pos) |
| 111 | possible_negs = list(set(self.queries.keys()) - set(neighbors)) |
| 112 | random.shuffle(possible_negs) |
| 113 | |
| 114 | query = self.load_file_func(dict_value["query"]) # Nx3 |
| 115 | query = np.expand_dims(query, axis=0) |
| 116 | positives = self.load_files_func(pos_files) |
| 117 | negatives = self.load_files_func(neg_files) |
| 118 | |
| 119 | output = [query, positives, negatives] |
| 120 | if other_neg: |
| 121 | neg2 = self.load_file_func(self.queries[possible_negs[0]]["query"]) |
| 122 | neg2 = np.expand_dims(neg2, axis=0) |
| 123 | output.append(neg2) |
| 124 | return output |
| 125 | |
| 126 | @abstractmethod |
no test coverage detected