Text feature wrapper
| 63 | return self.model.get_image_features(x) |
| 64 | |
| 65 | class TextFeatureExtractor(torch.nn.Module): |
| 66 | """ Text feature wrapper """ |
| 67 | def __init__(self, model): |
| 68 | super(TextFeatureExtractor, self).__init__() |
| 69 | self.model = model |
| 70 | |
| 71 | def __call__(self, x): |
| 72 | return self.model.get_text_features(x) |
| 73 | |
| 74 | def image_transform(t, height=7, width=7): |
| 75 | """ Transformation for CAM (image) """ |