From e583f3bcfce7242c43523ef4d27265105d3b7360 Mon Sep 17 00:00:00 2001 From: ryan-qiyu-jiang Date: Tue, 23 Nov 2021 14:56:03 -0800 Subject: [PATCH 1/4] [feat] Add pytorchvideo encoder wrapper Add an encoder class that constructs any pytorchvideo model from config, and uses this model for its forward pass. Can load pretrained or random init models, based on config. [ghstack-poisoned] --- mmf/modules/encoders.py | 89 +++++++++++++++++++++++++++++++++- requirements.txt | 2 + tests/modules/test_encoders.py | 33 ++++++++++++- tests/test_utils.py | 7 +++ 4 files changed, 128 insertions(+), 3 deletions(-) diff --git a/mmf/modules/encoders.py b/mmf/modules/encoders.py index 8a67bb1cc..191178a7f 100644 --- a/mmf/modules/encoders.py +++ b/mmf/modules/encoders.py @@ -1,10 +1,13 @@ # Copyright (c) Facebook, Inc. and its affiliates. +import importlib +import inspect +import logging import os import pickle import re from collections import OrderedDict from copy import deepcopy -from dataclasses import dataclass +from dataclasses import asdict, dataclass from enum import Enum from typing import Any @@ -25,13 +28,15 @@ from transformers.configuration_auto import AutoConfig from transformers.modeling_auto import AutoModel - try: from detectron2.modeling import ShapeSpec, build_resnet_backbone except ImportError: pass +logger = logging.getLogger() + + class Encoder(nn.Module): @dataclass class Config: @@ -688,6 +693,86 @@ def forward(self, x: Tensor) -> Tensor: return out +@registry.register_encoder("torchvideo") +class TorchVideoEncoder(Encoder): + """ + Wrapper around importing torchvideo models + as encoders. + """ + + @dataclass + class Config(Encoder.Config): + name: str = "torchvideo" + random_init: bool = False + model_name: str = "slowfast_r50" + cls_layer_num: int = 1 + + def __init__(self, config: Config): + pytorchvideo_spec = importlib.util.find_spec("pytorchvideo") + if pytorchvideo_spec is None: + raise ImportError("pytorchvideo required for using TorchVideoEncoder") + import pytorchvideo.models as models + + super().__init__() + config = OmegaConf.create({**asdict(self.Config()), **config}) + if config.random_init: + model_create_fn_name = f"create_{config.model_name}" + model_create_fn = getattr(models, model_create_fn_name) + params = dict(**config) + params.pop("random_init") + params.pop("model_name") + params.pop("cls_layer_num") + + accepted_params, ignored_params = self.filter_dict_to_signature( + model_create_fn, params + ) + if ignored_params: + ignored_params_str = " ".join(ignored_params.keys()) + logger.warning( + "The following model constructor params were ignored" + + " because they don't match a named param in the constructor: " + + ignored_params_str + ) + model = model_create_fn(**accepted_params) + else: + # load weights from TorchHub + model = torch.hub.load( + "facebookresearch/pytorchvideo:main", + model=config.model_name, + pretrained=True, + ) + + if config.cls_layer_num == 0: + self.encoder = model + return + + modules_list = list(model.children()) + if len(modules_list) == 1: + modules_list = list(modules_list[0].children()) + modules = modules_list[: -config.cls_layer_num] + self.encoder = nn.Sequential(*modules) + + def forward(self, *args, **kwargs): + # pass along input to model + # assumes caller obeys the dynamic model signature + return self.encoder(*args, **kwargs) + + def filter_dict_to_signature(self, callable, params): + constructor_signature = inspect.signature(callable) # Signature obj + constructor_params = constructor_signature.parameters + accepted_params = { + param_name: params[param_name] + for param_name in params + if param_name in constructor_params + } + ignored_params = { + param_name: params[param_name] + for param_name in params + if param_name not in constructor_params + } + return accepted_params, ignored_params + + @registry.register_encoder("r2plus1d_18") class R2Plus1D18VideoEncoder(PooledEncoder): """ diff --git a/requirements.txt b/requirements.txt index 2f0317749..ad045cb9a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -22,3 +22,5 @@ pytorch-lightning @ git+https://github.com/PyTorchLightning/pytorch-lightning@fa torchaudio>=0.6.0, <=0.9.0 psutil pillow==8.3.1 +av>=8.0.3 +pytorchvideo>=0.1.3 diff --git a/tests/modules/test_encoders.py b/tests/modules/test_encoders.py index 25e8063c0..7a5f3bfdc 100644 --- a/tests/modules/test_encoders.py +++ b/tests/modules/test_encoders.py @@ -6,7 +6,11 @@ import torch from mmf.modules import encoders from omegaconf import OmegaConf -from tests.test_utils import setup_proxy, skip_if_old_transformers +from tests.test_utils import ( + setup_proxy, + skip_if_old_transformers, + skip_if_no_pytorchvideo, +) from torch import nn @@ -102,3 +106,30 @@ def test_vit_encoder(self): x = torch.rand(32, 197, 768) output, _ = encoder(x) self.assertEqual(output.size(-1), config.out_dim) + + @skip_if_no_pytorchvideo + def test_torchvision_slowfast_r50_encoder(self): + config = OmegaConf.structured(encoders.TorchVideoEncoder.Config()) + encoder = encoders.TorchVideoEncoder(config) + fast = torch.rand((1, 3, 32, 224, 224)) + slow = torch.rand((1, 3, 8, 224, 224)) + output = encoder([slow, fast]) + self.assertEqual(output.size(1), 2304) + + @skip_if_no_pytorchvideo + def test_torchvision_mvit_encoder(self): + config = OmegaConf.create( + { + "name": "torchvideo", + "model_name": "multiscale_vision_transformers", + "random_init": True, + "cls_layer_num": 0, + "spatial_size": 224, + "temporal_size": 8, + "head": None, + } + ) + encoder = encoders.TorchVideoEncoder(config) + x = torch.rand((1, 3, 8, 224, 224)) + output = encoder(x) + self.assertEqual(output.shape, torch.Size([1, 12545, 96])) diff --git a/tests/test_utils.py b/tests/test_utils.py index c35433164..31df0b9b4 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -102,6 +102,13 @@ def wrap(testfn, reason="Requires newer version of transformers"): return wrap +def skip_if_no_pytorchvideo(testfn, reason="Requires pytorchvideo"): + import importlib + + pytorchvideo_spec = importlib.util.find_spec("pytorchvideo") + return unittest.skipUnless(pytorchvideo_spec is not None, reason)(testfn) + + def compare_state_dicts(a, b): same = True same = same and (list(a.keys()) == list(b.keys())) From f1eb042dbf3a433f6caa52339559ff7f12432a22 Mon Sep 17 00:00:00 2001 From: ryan-qiyu-jiang Date: Wed, 1 Dec 2021 13:39:33 -0800 Subject: [PATCH 2/4] Update on "[feat] Add pytorchvideo encoder wrapper" Add an encoder class that constructs any pytorchvideo model from config, and uses this model for its forward pass. Can load pretrained or random init models, based on config. Differential Revision: [D32631207](https://our.internmc.facebook.com/intern/diff/D32631207) [ghstack-poisoned] --- mmf/modules/encoders.py | 27 +++++++++++++- tests/models/test_mmf_transformer.py | 55 +++++++++++++++++++++++++++- tests/modules/test_encoders.py | 8 ++-- 3 files changed, 84 insertions(+), 6 deletions(-) diff --git a/mmf/modules/encoders.py b/mmf/modules/encoders.py index 191178a7f..cebbb0a9e 100644 --- a/mmf/modules/encoders.py +++ b/mmf/modules/encoders.py @@ -9,7 +9,7 @@ from copy import deepcopy from dataclasses import asdict, dataclass from enum import Enum -from typing import Any +from typing import Any, Optional import torch import torchvision @@ -773,6 +773,31 @@ def filter_dict_to_signature(self, callable, params): return accepted_params, ignored_params +@registry.register_encoder("mvit") +class MViTEncoder(Encoder): + """ + MVIT from pytorchvideo + """ + + @dataclass + class Config(Encoder.Config): + name: str = "mvit" + random_init: bool = False + model_name: str = "multiscale_vision_transformers" + spatial_size: int = 224 + temporal_size: int = 8 + head: Optional[Any] = None + + def __init__(self, config: Config): + super().__init__() + self.encoder = TorchVideoEncoder(config) + + def forward(self, *args, **kwargs): + output = self.encoder(*args, **kwargs) + output = output.permute(0, 2, 1) + return output[:, :1, :] + + @registry.register_encoder("r2plus1d_18") class R2Plus1D18VideoEncoder(PooledEncoder): """ diff --git a/tests/models/test_mmf_transformer.py b/tests/models/test_mmf_transformer.py index 63f0e0259..45074af19 100644 --- a/tests/models/test_mmf_transformer.py +++ b/tests/models/test_mmf_transformer.py @@ -21,7 +21,9 @@ from mmf.utils.configuration import Configuration from mmf.utils.env import setup_imports, teardown_imports from omegaconf import OmegaConf - +from tests.test_utils import ( + skip_if_no_pytorchvideo, +) BERT_VOCAB_SIZE = 30255 ROBERTA_VOCAB_SIZE = 50265 @@ -444,6 +446,57 @@ def test_preprocessing_with_resnet_encoder(self): test_utils.compare_tensors(segment_ids["image"], torch.tensor([[0], [0]])) test_utils.compare_tensors(segment_ids["text"], torch.ones((2, 128)).long()) + @skip_if_no_pytorchvideo + def test_preprocessing_with_mvit_encoder(self): + encoder_config = OmegaConf.create( + { + "name": "mvit", + "model_name": "multiscale_vision_transformers", + "random_init": True, + "cls_layer_num": 0, + "spatial_size": 224, + "temporal_size": 8, + "head": None, + } + ) + self._image_modality_config = MMFTransformerModalityConfig( + type="image", + key="image", + embedding_dim=12545, + position_dim=1, + segment_id=0, + encoder=encoder_config, + ) + modalities_config = [self._image_modality_config, self._text_modality_config] + config = MMFTransformer.Config(modalities=modalities_config, num_labels=2) + mmft = build_model(config) + + sample_list = SampleList() + sample_list.image = torch.rand((2, 3, 8, 224, 224)) + sample_list.text = torch.randint(0, 512, (2, 128)) + + transformer_input = mmft.preprocess_sample(sample_list) + input_ids = transformer_input["input_ids"] + self.assertEqual(input_ids["image"].dim(), 3) + self.assertEqual(list(input_ids["image"].size()), [2, 1, 12545]) + + self.assertEqual(input_ids["text"].dim(), 2) + self.assertEqual(list(input_ids["text"].size()), [2, 128]) + + position_ids = transformer_input["position_ids"] + test_utils.compare_tensors(position_ids["image"], torch.tensor([[0], [0]])) + test_utils.compare_tensors( + position_ids["text"], torch.arange(0, 128).unsqueeze(0).expand((2, 128)) + ) + + masks = transformer_input["masks"] + test_utils.compare_tensors(masks["image"], torch.tensor([[1], [1]])) + test_utils.compare_tensors(masks["text"], torch.ones((2, 128)).long()) + + segment_ids = transformer_input["segment_ids"] + test_utils.compare_tensors(segment_ids["image"], torch.tensor([[0], [0]])) + test_utils.compare_tensors(segment_ids["text"], torch.ones((2, 128)).long()) + def test_tie_mlm_head_weight_to_encoder(self): self._text_modality_config = MMFTransformerModalityConfig( type="text", diff --git a/tests/modules/test_encoders.py b/tests/modules/test_encoders.py index 7a5f3bfdc..7c185e962 100644 --- a/tests/modules/test_encoders.py +++ b/tests/modules/test_encoders.py @@ -117,10 +117,10 @@ def test_torchvision_slowfast_r50_encoder(self): self.assertEqual(output.size(1), 2304) @skip_if_no_pytorchvideo - def test_torchvision_mvit_encoder(self): + def test_mvit_encoder(self): config = OmegaConf.create( { - "name": "torchvideo", + "name": "mvit", "model_name": "multiscale_vision_transformers", "random_init": True, "cls_layer_num": 0, @@ -129,7 +129,7 @@ def test_torchvision_mvit_encoder(self): "head": None, } ) - encoder = encoders.TorchVideoEncoder(config) + encoder = encoders.MViTEncoder(config) x = torch.rand((1, 3, 8, 224, 224)) output = encoder(x) - self.assertEqual(output.shape, torch.Size([1, 12545, 96])) + self.assertEqual(output.shape, torch.Size([1, 1, 12545])) From bbe64ca70609259b6a35b552aa9c8cc5e1abbc0e Mon Sep 17 00:00:00 2001 From: ryan-qiyu-jiang Date: Thu, 2 Dec 2021 09:34:31 -0800 Subject: [PATCH 3/4] Update on "[feat] Add pytorchvideo encoder wrapper" Add an encoder class that constructs any pytorchvideo model from config, and uses this model for its forward pass. Can load pretrained or random init models, based on config. Differential Revision: [D32631207](https://our.internmc.facebook.com/intern/diff/D32631207) [ghstack-poisoned] --- mmf/modules/encoders.py | 46 +++++++++++++++++++++++++--- tests/models/test_mmf_transformer.py | 2 +- tests/modules/test_encoders.py | 44 ++++++++++++++++++-------- 3 files changed, 73 insertions(+), 19 deletions(-) diff --git a/mmf/modules/encoders.py b/mmf/modules/encoders.py index cebbb0a9e..d60ceea93 100644 --- a/mmf/modules/encoders.py +++ b/mmf/modules/encoders.py @@ -9,7 +9,7 @@ from copy import deepcopy from dataclasses import asdict, dataclass from enum import Enum -from typing import Any, Optional +from typing import Any, Optional, List import torch import torchvision @@ -718,7 +718,7 @@ def __init__(self, config: Config): if config.random_init: model_create_fn_name = f"create_{config.model_name}" model_create_fn = getattr(models, model_create_fn_name) - params = dict(**config) + params = dict(**OmegaConf.to_container(config)) params.pop("random_init") params.pop("model_name") params.pop("cls_layer_num") @@ -786,16 +786,52 @@ class Config(Encoder.Config): model_name: str = "multiscale_vision_transformers" spatial_size: int = 224 temporal_size: int = 8 + encoder_pool_type: str = "cls" head: Optional[Any] = None + embed_dim_mul: Optional[List] = None + atten_head_mul: Optional[List] = None + pool_q_stride_size: Optional[List] = None + pool_kv_stride_adaptive: Optional[List] = None + pool_kvq_kernel: Optional[List] = None def __init__(self, config: Config): super().__init__() - self.encoder = TorchVideoEncoder(config) + config = {**asdict(self.Config()), **config} + # initialize default lists + config["embed_dim_mul"] = config["embed_dim_mul"] or [ + [1, 2.0], + [3, 2.0], + [14, 2.0], + ] + config["atten_head_mul"] = config["atten_head_mul"] or [ + [1, 2.0], + [3, 2.0], + [14, 2.0], + ] + config["pool_q_stride_size"] = config["pool_q_stride_size"] or [ + [1, 1, 2, 2], + [3, 1, 2, 2], + [14, 1, 2, 2], + ] + config["pool_kv_stride_adaptive"] = config["pool_kv_stride_adaptive"] or [ + 1, + 8, + 8, + ] + config["pool_kvq_kernel"] = config["pool_kvq_kernel"] or [3, 3, 3] + + self.pool_type = config.pop("encoder_pool_type") + self.encoder = TorchVideoEncoder(OmegaConf.create(config)) def forward(self, *args, **kwargs): output = self.encoder(*args, **kwargs) - output = output.permute(0, 2, 1) - return output[:, :1, :] + if self.pool_type == "cls": + output = output[:, :1, :] + elif self.pool_type == "avg": + output = output.mean(1).unsqueeze(1) + elif self.pool_type == "identity": + output = output + return output @registry.register_encoder("r2plus1d_18") diff --git a/tests/models/test_mmf_transformer.py b/tests/models/test_mmf_transformer.py index 45074af19..ec222513c 100644 --- a/tests/models/test_mmf_transformer.py +++ b/tests/models/test_mmf_transformer.py @@ -478,7 +478,7 @@ def test_preprocessing_with_mvit_encoder(self): transformer_input = mmft.preprocess_sample(sample_list) input_ids = transformer_input["input_ids"] self.assertEqual(input_ids["image"].dim(), 3) - self.assertEqual(list(input_ids["image"].size()), [2, 1, 12545]) + self.assertEqual(list(input_ids["image"].size()), [2, 1, 768]) self.assertEqual(input_ids["text"].dim(), 2) self.assertEqual(list(input_ids["text"].size()), [2, 128]) diff --git a/tests/modules/test_encoders.py b/tests/modules/test_encoders.py index 7c185e962..c6b5ec08c 100644 --- a/tests/modules/test_encoders.py +++ b/tests/modules/test_encoders.py @@ -118,18 +118,36 @@ def test_torchvision_slowfast_r50_encoder(self): @skip_if_no_pytorchvideo def test_mvit_encoder(self): - config = OmegaConf.create( - { - "name": "mvit", - "model_name": "multiscale_vision_transformers", - "random_init": True, - "cls_layer_num": 0, - "spatial_size": 224, - "temporal_size": 8, - "head": None, - } - ) - encoder = encoders.MViTEncoder(config) + config = { + "name": "mvit", + "model_name": "multiscale_vision_transformers", + "random_init": True, + "cls_layer_num": 0, + "spatial_size": 224, + "temporal_size": 8, + "head": None, + "embed_dim_mul": [[1, 2.0], [3, 2.0], [14, 2.0]], + "atten_head_mul": [[1, 2.0], [3, 2.0], [14, 2.0]], + "pool_q_stride_size": [[1, 1, 2, 2], [3, 1, 2, 2], [14, 1, 2, 2]], + "pool_kv_stride_adaptive": [1, 8, 8], + "pool_kvq_kernel": [3, 3, 3], + } + # test bert cls pooler + encoder = encoders.MViTEncoder(OmegaConf.create(config)) x = torch.rand((1, 3, 8, 224, 224)) output = encoder(x) - self.assertEqual(output.shape, torch.Size([1, 1, 12545])) + self.assertEqual(output.shape, torch.Size([1, 1, 768])) + + # test avg pooler + encoder = encoders.MViTEncoder( + OmegaConf.create(dict(config, encoder_pool_type="avg")) + ) + output = encoder(x) + self.assertEqual(output.shape, torch.Size([1, 1, 768])) + + # test no pooling + encoder = encoders.MViTEncoder( + OmegaConf.create(dict(config, encoder_pool_type="identity")) + ) + output = encoder(x) + self.assertEqual(output.shape, torch.Size([1, 197, 768])) From 434a12abe789533857d155e2682c94f102aac38e Mon Sep 17 00:00:00 2001 From: ryan-qiyu-jiang Date: Thu, 2 Dec 2021 14:53:54 -0800 Subject: [PATCH 4/4] Update on "[feat] Add pytorchvideo encoder wrapper" Add an encoder class that constructs any pytorchvideo model from config, and uses this model for its forward pass. Can load pretrained or random init models, based on config. Differential Revision: [D32631207](https://our.internmc.facebook.com/intern/diff/D32631207) [ghstack-poisoned] --- mmf/modules/encoders.py | 13 ++++--------- tests/modules/test_encoders.py | 12 +++++++++++- tests/test_utils.py | 2 +- 3 files changed, 16 insertions(+), 11 deletions(-) diff --git a/mmf/modules/encoders.py b/mmf/modules/encoders.py index d60ceea93..fab85ffaa 100644 --- a/mmf/modules/encoders.py +++ b/mmf/modules/encoders.py @@ -695,10 +695,7 @@ def forward(self, x: Tensor) -> Tensor: @registry.register_encoder("torchvideo") class TorchVideoEncoder(Encoder): - """ - Wrapper around importing torchvideo models - as encoders. - """ + """Wrapper around importing torchvideo models""" @dataclass class Config(Encoder.Config): @@ -730,8 +727,8 @@ def __init__(self, config: Config): ignored_params_str = " ".join(ignored_params.keys()) logger.warning( "The following model constructor params were ignored" - + " because they don't match a named param in the constructor: " - + ignored_params_str + + " because they don't match a named param in the" + + f" constructor: {ignored_params_str}" ) model = model_create_fn(**accepted_params) else: @@ -775,9 +772,7 @@ def filter_dict_to_signature(self, callable, params): @registry.register_encoder("mvit") class MViTEncoder(Encoder): - """ - MVIT from pytorchvideo - """ + """MVIT from pytorchvideo""" @dataclass class Config(Encoder.Config): diff --git a/tests/modules/test_encoders.py b/tests/modules/test_encoders.py index c6b5ec08c..09b3b380a 100644 --- a/tests/modules/test_encoders.py +++ b/tests/modules/test_encoders.py @@ -108,12 +108,16 @@ def test_vit_encoder(self): self.assertEqual(output.size(-1), config.out_dim) @skip_if_no_pytorchvideo - def test_torchvision_slowfast_r50_encoder(self): + def test_torchvideo_slowfast_r50_encoder(self): + # instantiate video encoder from pytorchvideo + # default model is slowfast_r50 config = OmegaConf.structured(encoders.TorchVideoEncoder.Config()) encoder = encoders.TorchVideoEncoder(config) fast = torch.rand((1, 3, 32, 224, 224)) slow = torch.rand((1, 3, 8, 224, 224)) output = encoder([slow, fast]) + # check output tensor is the expected feature dim size + # (bs, feature_dim) self.assertEqual(output.size(1), 2304) @skip_if_no_pytorchvideo @@ -136,6 +140,11 @@ def test_mvit_encoder(self): encoder = encoders.MViTEncoder(OmegaConf.create(config)) x = torch.rand((1, 3, 8, 224, 224)) output = encoder(x) + # check output tensor is the expected feature dim size + # based on pooled attention configs + # for more details consult https://arxiv.org/pdf/2104.11227 + # and https://github.com/facebookresearch/pytorchvideo/ + # (bs, num_features, feature_dim) self.assertEqual(output.shape, torch.Size([1, 1, 768])) # test avg pooler @@ -150,4 +159,5 @@ def test_mvit_encoder(self): OmegaConf.create(dict(config, encoder_pool_type="identity")) ) output = encoder(x) + # (bs, num_features, feature_dim) self.assertEqual(output.shape, torch.Size([1, 197, 768])) diff --git a/tests/test_utils.py b/tests/test_utils.py index 31df0b9b4..d48558543 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -106,7 +106,7 @@ def skip_if_no_pytorchvideo(testfn, reason="Requires pytorchvideo"): import importlib pytorchvideo_spec = importlib.util.find_spec("pytorchvideo") - return unittest.skipUnless(pytorchvideo_spec is not None, reason)(testfn) + return unittest.skipIf(pytorchvideo_spec is None, reason)(testfn) def compare_state_dicts(a, b):