SCENIC: A JAX Library for Computer Vision Research and Beyond

Mostafa Dehghani, Alexey Gritsenko, Anurag Arnab, Matthias Minderer, Yi Tay

Introduction

It is an exciting time for architecture research in computer vision. With new architectures like ViT taking the world by storm, there exists a clear demand for software and machine learning infrastructure to support easy and extensible neural network architecture research in the field of vision. As attention models and MLP-only architectures become more popular, we expect to see even more research in the coming years pushing the field forward.

We introduce Scenic, an open-source JAX library for fast and extensible research in vision and beyond. Scenic has been successfully used to develop classification, segmentation, and detection models for images, videos, and other modalities, including multi-modal setups.

Scenic strives to be a unified, all-in-one codebase for modeling needs, currently offering implementations of state-of-the-art vision models like ViT , DETR , MLP Mixer , ResNet , and U-Net . On top of that, Scenic has been used for numerous Google projects and research papers such as ViViT , OmniNet , TokenLearner , MBT , studies on scaling behaviour of various models , and others. We anticipate more research projects to be open-sourced in the Scenic repository in the near future.

Scenic is developed in JAX and uses Flax as the neural network library, relies on TFDS and DMVR for implementing the input pipeline of most of the tasks and dataset, and makes use of for common training loop functionalities offered by CLU . JAX is an “ultra-simple to use” library that enables automatic differentiation of native Python and NumPy functions. Moreover, it supports multi-host and multi-device training on accelerators including GPUs and TPUs, making it ideal for large-scale machine learning research.

In a nutshell, Scenic is a (i) set of shared light-weight libraries solving commonly encountered tasks when training large-scale (i.e. multi-device, multi-host) models in vision and beyond; and (ii) a number of projects containing fully fleshed out problem-specific training and evaluation loops using these libraries.

Scenic is designed to propose different levels of abstraction. It supports projects from those that only require changing hyper-parameters, to those that need customization on the input pipeline, model architecture, losses and metrics, and the training loop. To make this happen, the code in Scenic is organized as either project-level code, which refers to customized code for specific projects or baselines, or library-level code, which refers to common functionalities and general patterns that are adapted by the majority of projects. The project-level code lives in the projects directory.

Scenic aims to facilitate the rapid prototyping of large-scale models. To keep the code simple to understand and extend, Scenic design prefers forking and copy-pasting over adding complexity or increasing abstraction. We only upstream functionality to the library-level when it proves to be widely useful across multiple models and tasks. Minimizing support for various use-cases in the library-level code helps us to avoid accumulating generalizations that result in the code being complex and difficult to understand. Note that complexity or abstractions of any level can be added to project-level code.

Design

Scenic offers a unified framework that is sufficiently flexible to support projects in a wide range of needs without having to write complex code. Scenic contains optimized implementations of a set of research models operating on a wide range of modalities (video, image, audio, and text), and supports several datasets. This again is made possible by its flexible and low-overhead design. In this section, we go over different parts and discuss the structure that is used to organize projects and library code.

The goal is to keep the library-level code minimal and well-tested and to avoid introducing extra abstractions to support minor use-cases. Shared libraries provided by Scenic are split into:

dataset_lib: Implements IO pipelines for loading and pre-processing data for common tasks and benchmarks. All pipelines are designed to be scalable and support multi-host and multi-device setups, taking care of dividing data among multiple hosts, incomplete batches, caching, pre-fetching, etc.

model_lib : Provides several abstract model interfaces (e.g., ClassificationModel or SegmentationModel in model_lib/base_models) with task-specific losses and metrics; neural network layers in model_lib/layers, focusing on efficient implementation of attention and transformer primitives; and finally accelerator-friendly implementations of bipartite matching algorithms in model_lib/matchers.

train_lib: Provides tools for constructing training loops and implements several optimized trainers (classification trainer and segmentation trainer) that can be forked for customization.

common_lib: General utilities, such as logging and debugging modules, and functionalities for processing raw data.

2 Project-level code

Scenic supports the development of customized solutions for specialized tasks and data via the concept of the “project”. There is no one-fits-all recipe for how much code should be re-used by a project. Projects can consist of only configuration files and use the common models, trainers, tasks/data that live in library-level code, or they can simply fork any of the mentioned functionalities and redefine, layers, losses, metrics, logging methods, tasks, architectures, as well as training and evaluation loops. The modularity of library-level code makes it flexible enough to support projects falling anywhere on the “run-as-is” to “fully-customized” spectrum.

Common baselines such as a ResNet, Vision Transformer (ViT), and DETR are implemented in the projects/baselines project. Forking models in this directory is a good starting point for new projects.

3 Scenic BaseModel

A solution usually has several parts: data/task pipeline, model architecture, losses and metrics, training and evaluation, etc. Given that much of the research done in Scenic is trying out different architectures, Scenic introduces the concept of a “model”, to facilitate “plug-in/plug-out” experiments. A Scenic model is defined as the network architecture, the losses that are used to update the weights of the network during training, and metrics that are used to evaluate the output of the network. This is implemented as BaseModel.

BaseModel is an abstract class with three members: a build_flax_model, a loss_fn, and a get_metrics_fn.

build_flax_model function returns a flax_model. A typical usage pattern is depicted below:

Abstract classes for defining Scenic models are declared in model_lib/base_models. These include the BaseModel that all models inherit from, as well as ClassificationModel, MultiLabelClassificationModel, EncoderDecoderModel and SegmentationModel that respectively define losses and metrics for classification, seq2seq, and segmentation tasks. Depending on its needs, a Scenic project can define new base class or override an existing one for its specific tasks, losses and metrics.

A typical model loss function in Scenic expects predictions and a batch of data:

Finally, a typical get_metrics_fn returns a callable, metric_fn, that calculates appropriate metrics and returns them as a Python dictionary. The metric function, for each metric, computes f(xi,yi)f(x_{i},y_{i}) on a mini-batch, where xix_{i} and yiy_{i} are inputs and labels of iith example, and returns a dictionary from the metric name to a tuple of metric value and the metric normalizer (typically, the number of examples in the mini-batch). It has the API:

Given metric and normalizer values collected from examples processed by all devices in all hosts, the model trainer is then responsible for aggregating and computing the normalized metric value for the evaluated examples.

Importantly, while the design pattern above is recommended and has been found to work well for a range of projects, it is not forced, and there is no issue deviating from the above structure within a project.

Conclusion

Machine Learning (ML) infrastructure is a cornerstone of ML research. Enabling researchers to quickly try out new ideas, and to rapidly scale them up when they show promise, accelerates research. Furthermore, history suggests that methods that leverage the computation available at the time are often the most effective . Scenic embodies our experience of developing the best research infrastructure, and we are excited to share it with the broader community. We hope to see many more brilliant ideas being developed using Scenic, contributing to the amazing progress made by the ML community for improving lives through AI.

Acknowledgment

Scenic has been developed and improved with the help of many amazing contributors. We would like to especially thank Dirk Weissenborn, Andreas Steiner, Marvin Ritter, Aravindh Mahendran, Samira Abnar, Sunayana Rane, Josip Djolonga, Lucas Beyer, Alexander Kolesnikov, Xiaohua Zhai, Rob Romijnders, Rianne van den Berg, Jonathan Heek, Olivier Teboul, Marco Cuturi, Lu Jiang, Mario Lučić, and Neil Houlsby for their direct or indirect contributions to Scenic.

References