Image feature wrapper
| 54 | return cos(model_output, self.features) |
| 55 | |
| 56 | class ImageFeatureExtractor(torch.nn.Module): |
| 57 | """ Image feature wrapper """ |
| 58 | def __init__(self, model): |
| 59 | super(ImageFeatureExtractor, self).__init__() |
| 60 | self.model = model |
| 61 | |
| 62 | def __call__(self, x): |
| 63 | return self.model.get_image_features(x) |
| 64 | |
| 65 | class TextFeatureExtractor(torch.nn.Module): |
| 66 | """ Text feature wrapper """ |