class VideoMAE(nn.Module):
def __init__(self, layer=11, **kwargs):
super().__init__()
try:
from transformers import VideoMAEForVideoClassification
except ImportError:
raise ImportError(
"Please install the transformers library: pip install transformers"
)
self.model = VideoMAEForVideoClassification.from_pretrained(
"MCG-NJU/videomae-base-finetuned-kinetics"
)
self.model.requires_grad_(False)
self.model.eval()
self.layer = layer
def forward(self, x):
assert x.dim() == 5
assert x.shape[1:] == (16, 3, 224, 224) # frame, color channel, height, width
outputs = self.model(x, output_hidden_states=True, return_dict=True)
layer_idx = - (12 - self.layer)
return outputs.hidden_states[layer_idx]
# last_layer = outputs.hidden_states[-1]
# return last_layer
from PIL import Image
from torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize
import numpy as np
def transform_images(frames, size=(224, 224)):
resized = []
length = len(frames)
for i in range(length):
frame = frames[i]
# image = Image.fromarray((frame * 255).astype(np.uint8))
image = Image.fromarray(frame)
image = image.resize(size, Image.ANTIALIAS)
image = np.array(image) / 255.0
resized.append(np.array(image))
frames = np.stack(resized, axis=0)
frames = frames.transpose(0, 3, 1, 2) # (N, H, W, C) -> (N, C, H, W)
frames = torch.tensor(frames, dtype=torch.float32)
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
frames = (frames - mean) / std
return frames
def read_video(video_path: str) -> torch.Tensor:
try:
from decord import VideoReader
except ImportError:
raise ImportError("Please install the decord library: pip install decord")
vr = VideoReader(video_path)
print(f"Total frames: {len(vr)}")
# frames = vr.get_batch(range(len(vr))).asnumpy()
lenth = len(vr)
lenth = 1600 if lenth > 1600 else lenth
frames = vr.get_batch(np.arange(lenth)).asnumpy()
# if less than 1600 frames, repeat the last frame
if lenth < 1600:
last_frame = frames[-1]
for i in range(1600 - lenth):
frames = np.append(frames, last_frame.reshape(1, *last_frame.shape), axis=0)
# frames = np.array(frames)
frames = transform_images(frames)
return frames
def video_mae_feature(video_path, layer=11):
frames = read_video(video_path)
videomae = VideoMAE(layer=layer)
videomae = videomae.cuda()
frames = frames.cuda()
frames = rearrange(frames, "(b t) c h w -> b t c h w", t=16)
feats = videomae(frames)
return feats # (t/2, (h*w), c)