(self, ids, masks, layers, interventions)
| 149 | return [dict(d=d) for d in imgurls] |
| 150 | |
| 151 | def get_features(self, ids, masks, layers, interventions): |
| 152 | zs = self.get_zs_for_ids(ids) |
| 153 | z_tensor = torch.tensor(zs).float().to(self.tester.device) |
| 154 | t_masks = torch.stack( |
| 155 | [torch.from_numpy(mask_to_numpy(mask)) for mask in masks] |
| 156 | )[:,None,:,:].to(self.tester.device) |
| 157 | t_features = self.tester.feature_stats(z_tensor, t_masks, |
| 158 | decode_intervention_array(interventions, |
| 159 | self.tester.layer_shapes()), layers) |
| 160 | # Convert torch arrays to plain python lists before returning. |
| 161 | return { layer: { key: value.cpu().numpy().tolist() |
| 162 | for key, value in feature.items() } |
| 163 | for layer, feature in t_features.items() } |
| 164 | |
| 165 | def get_featuremaps(self, ids, layers, interventions): |
| 166 | zs = self.get_zs_for_ids(ids) |
no test coverage detected