MCPcopy Create free account
hub / github.com/CSAILVision/gandissect / get_features

Method get_features

netdissect/serverstate.py:151–163  ·  view source on GitHub ↗
(self, ids, masks, layers, interventions)

Source from the content-addressed store, hash-verified

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)

Callers 1

post_featuresFunction · 0.80

Calls 5

get_zs_for_idsMethod · 0.95
mask_to_numpyFunction · 0.85
feature_statsMethod · 0.80
layer_shapesMethod · 0.80

Tested by

no test coverage detected