Team Ai
Apppublic

pytorch/SlowFast

sourceHugging Faceupdated 5y agoView on Hugging Face
1likes
app.py131 linesDownload Raw Back to root
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)