pytorch/SlowFast
1
1import torch2# Choose the `slowfast_r50` model 3model = torch.hub.load('facebookresearch/pytorchvideo', 'slowfast_r50', pretrained=True)4from typing import Dict5import json6import urllib7from torchvision.transforms import Compose, Lambda8from torchvision.transforms._transforms_video import (9 CenterCropVideo,10 NormalizeVideo,11)12from pytorchvideo.data.encoded_video import EncodedVideo13from pytorchvideo.transforms import (14 ApplyTransformToKey,15 ShortSideScale,16 UniformTemporalSubsample,17 UniformCropVideo18) 19 20import gradio as gr21# Set to GPU or CPU22device = "cpu"23model = model.eval()24model = model.to(device)25json_url = "https://dl.fbaipublicfiles.com/pyslowfast/dataset/class_names/kinetics_classnames.json"26json_filename = "kinetics_classnames.json"27try: urllib.URLopener().retrieve(json_url, json_filename)28except: urllib.request.urlretrieve(json_url, json_filename)29with open(json_filename, "r") as f:30 kinetics_classnames = json.load(f)31 32# Create an id to label name mapping33kinetics_id_to_classname = {}34for k, v in kinetics_classnames.items():35 kinetics_id_to_classname[v] = str(k).replace('"', "")36side_size = 25637mean = [0.45, 0.45, 0.45]38std = [0.225, 0.225, 0.225]39crop_size = 25640num_frames = 3241sampling_rate = 242frames_per_second = 3043slowfast_alpha = 444num_clips = 1045num_crops = 346 47class PackPathway(torch.nn.Module):48 """49 Transform for converting video frames as a list of tensors. 50 """51 def __init__(self):52 super().__init__()53 54 def forward(self, frames: torch.Tensor):55 fast_pathway = frames56 # Perform temporal sampling from the fast pathway.57 slow_pathway = torch.index_select(58 frames,59 1,60 torch.linspace(61 0, frames.shape[1] - 1, frames.shape[1] // slowfast_alpha62 ).long(),63 )64 frame_list = [slow_pathway, fast_pathway]65 return frame_list66 67transform = ApplyTransformToKey(68 key="video",69 transform=Compose(70 [71 UniformTemporalSubsample(num_frames),72 Lambda(lambda x: x/255.0),73 NormalizeVideo(mean, std),74 ShortSideScale(75 size=side_size76 ),77 CenterCropVideo(crop_size),78 PackPathway()79 ]80 ),81)82 83# The duration of the input clip is also specific to the model.84clip_duration = (num_frames * sampling_rate)/frames_per_second85url_link = "https://dl.fbaipublicfiles.com/pytorchvideo/projects/archery.mp4"86video_path = 'archery.mp4'87try: urllib.URLopener().retrieve(url_link, video_path)88except: urllib.request.urlretrieve(url_link, video_path)89# Select the duration of the clip to load by specifying the start and end duration90# The start_sec should correspond to where the action occurs in the video91 92def inference(in_vid):93 start_sec = 094 end_sec = start_sec + clip_duration95 96 # Initialize an EncodedVideo helper class and load the video97 video = EncodedVideo.from_path(in_vid)98 99 # Load the desired clip100 video_data = video.get_clip(start_sec=start_sec, end_sec=end_sec)101 102 # Apply a transform to normalize the video input103 video_data = transform(video_data)104 105 # Move the inputs to the desired device106 inputs = video_data["video"]107 inputs = [i.to(device)[None, ...] for i in inputs]108 # Pass the input clip through the model109 preds = model(inputs)110 111 # Get the predicted classes112 post_act = torch.nn.Softmax(dim=1)113 preds = post_act(preds)114 pred_classes = preds.topk(k=5).indices[0]115 116 # Map the predicted classes to the label names117 pred_class_names = [kinetics_id_to_classname[int(i)] for i in pred_classes]118 return "%s" % ", ".join(pred_class_names)119 120inputs = gr.inputs.Video(label="Input Video")121outputs = gr.outputs.Textbox(label="Top 5 predicted labels")122 123title = "SLOWFAST"124description = "demo for SLOWFAST, SlowFast networks pretrained on the Kinetics 400 dataset. To use it, simply upload your video, or click one of the examples to load them. Read more at the links below."125article = "<p style='text-align: center'><a href='https://arxiv.org/abs/1812.03982'>SlowFast Networks for Video Recognition</a> | <a href='https://github.com/facebookresearch/pytorchvideo'>Github Repo</a></p>"126 127examples = [128 ['archery.mp4']129]130 131gr.Interface(inference, inputs, outputs, title=title, description=description, article=article, examples=examples, analytics_enabled=False).launch(debug=True)